mirror of
https://github.com/TheoLeeCJ/SemIf-OpenJev.git
synced 2026-09-28 07:02:55 +08:00
Tolerate float roundoff in MLX evidence verification
This commit is contained in:
@@ -2,6 +2,7 @@
|
||||
import argparse
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import statistics
|
||||
from pathlib import Path
|
||||
|
||||
@@ -14,6 +15,24 @@ def load(path):
|
||||
return read_json(path)
|
||||
|
||||
|
||||
def assert_json_close(actual, expected, path="$"):
|
||||
"""Require identical JSON structure while tolerating float roundoff."""
|
||||
assert type(actual) is type(expected), f"{path}: type mismatch"
|
||||
if isinstance(actual, dict):
|
||||
assert actual.keys() == expected.keys(), f"{path}: key mismatch"
|
||||
for key in actual:
|
||||
assert_json_close(actual[key], expected[key], f"{path}.{key}")
|
||||
elif isinstance(actual, list):
|
||||
assert len(actual) == len(expected), f"{path}: length mismatch"
|
||||
for index, (left, right) in enumerate(zip(actual, expected)):
|
||||
assert_json_close(left, right, f"{path}[{index}]")
|
||||
elif isinstance(actual, float):
|
||||
assert math.isclose(actual, expected, rel_tol=1e-12, abs_tol=1e-12), (
|
||||
f"{path}: {actual!r} != {expected!r}")
|
||||
else:
|
||||
assert actual == expected, f"{path}: {actual!r} != {expected!r}"
|
||||
|
||||
|
||||
def verify(root):
|
||||
manifest, summary = load(root / "manifest.json"), load(root / "summary.json")
|
||||
checks = 0
|
||||
@@ -64,7 +83,7 @@ def verify(root):
|
||||
gold, rows = read(f"benchmarks/data/{name}.jsonl"), read(root / f"{name}.jsonl")
|
||||
validate(rows, gold)
|
||||
original = read(f"results/raw/predictions/direct-{name}.jsonl")
|
||||
assert report[name]["evaluation"] == evaluate.evaluate(gold, rows, comparison=original)
|
||||
assert_json_close(report[name]["evaluation"], evaluate.evaluate(gold, rows, comparison=original))
|
||||
assert report[name]["vs_published_torch"] == compare(original, rows)
|
||||
assert not report[name]["vs_published_torch"]["prompt_mismatches"]
|
||||
assert summary["quality"][name] == report[name]["evaluation"]["mean_family_balanced_accuracy"]
|
||||
|
||||
@@ -3,6 +3,7 @@ import gzip
|
||||
import hashlib
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -11,6 +12,12 @@ spec = importlib.util.spec_from_file_location(
|
||||
evidence = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(evidence)
|
||||
|
||||
benchmarks = Path(__file__).resolve().parents[1] / "benchmarks"
|
||||
sys.path.insert(0, str(benchmarks))
|
||||
verify_spec = importlib.util.spec_from_file_location("verify_mlx", benchmarks / "verify_mlx.py")
|
||||
verify = importlib.util.module_from_spec(verify_spec)
|
||||
verify_spec.loader.exec_module(verify)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("compressed", [False, True])
|
||||
def test_original_checksums_accept_plain_or_gzip(tmp_path, compressed):
|
||||
@@ -40,3 +47,14 @@ def test_rejects_ambiguous_or_corrupt_compressed_evidence(tmp_path):
|
||||
compressed.write_bytes(b"not gzip")
|
||||
with pytest.raises(gzip.BadGzipFile):
|
||||
evidence.read_bytes(path)
|
||||
|
||||
|
||||
def test_verifier_tolerates_only_float_roundoff():
|
||||
verify.assert_json_close(
|
||||
{"score": 0.1 + 0.2, "choices": ["yes", "no"]},
|
||||
{"score": 0.3, "choices": ["yes", "no"]},
|
||||
)
|
||||
with pytest.raises(AssertionError):
|
||||
verify.assert_json_close({"score": 0.31}, {"score": 0.3})
|
||||
with pytest.raises(AssertionError):
|
||||
verify.assert_json_close({"choices": ["no"]}, {"choices": ["yes"]})
|
||||
|
||||
Reference in New Issue
Block a user