Tolerate float roundoff in MLX evidence verification

This commit is contained in:
TheoLeeCJ
2026-09-19 12:46:33 +08:00
parent 7d107b05b8
commit ca3ba65f14
2 changed files with 38 additions and 1 deletions
+20 -1
View File
@@ -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"]
+18
View File
@@ -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"]})