diff --git a/benchmarks/bench.py b/benchmarks/bench.py index 74f19c5..ae13288 100644 --- a/benchmarks/bench.py +++ b/benchmarks/bench.py @@ -71,6 +71,17 @@ def number(value): ) +WIRE_DECIMALS = 4 # workers round probabilities on the wire to four decimal places + + +def sum_allowance(count): + """How far `count` wire probabilities may miss a normalized sum: each rounded value + carries up to half a decimal step, so the allowance grows with the label count. A + flat 1e-4 rejected four-value distributions summing to 0.9999 at the float boundary + (PR #40, first GPU run).""" + return count * 0.5 * 10**-WIRE_DECIMALS + 1e-9 + + def read_answers(case, payload): """Validate wire results and normalize probabilities by label, never by order.""" answers = payload["answers"] @@ -97,7 +108,8 @@ def read_answers(case, payload): value = answer[kind] if any(not number(p) or not 0 <= p <= 1 for p in probabilities.values()): raise ValueError("probabilities must be finite and in [0, 1]") - if not math.isclose(sum(probabilities.values()), 1, abs_tol=1e-4): + total = math.fsum(probabilities.values()) + if abs(total - 1) > sum_allowance(len(probabilities)): raise ValueError("probabilities must sum to one") if kind == "choice": if value not in probabilities or probabilities[value] != max( diff --git a/tests/benchmarks/test_bench.py b/tests/benchmarks/test_bench.py index 99cf8b0..3ca1791 100644 --- a/tests/benchmarks/test_bench.py +++ b/tests/benchmarks/test_bench.py @@ -176,6 +176,32 @@ def test_all_primitives_and_invalid_distributions(self): with self.assertRaises(ValueError): bench.read_answers(CASES[2], payload) + def test_rounded_probabilities_at_the_sum_boundary(self): + """PR #40's first GPU run: the worker rounds probabilities to four decimals, and + four rounded values legitimately miss 1 by a full rounding step — a flat 1e-4 + tolerance rejected both vectors below at the floating-point boundary.""" + case = { + "id": "mmlu-test-10312", + "request": { + "questions": {"q": {"type": "choice", "criteria": ["A", "B", "C", "D"]}} + }, + "expected": {}, + } + for probabilities in ( + {"A": 0.1099, "B": 0.1211, "C": 0.5589, "D": 0.2100}, + {"A": 0.1130, "B": 0.5135, "C": 0.3151, "D": 0.0583}, + ): + answer = { + "choice": max(probabilities, key=probabilities.get), + "probabilities": probabilities, + } + bench.read_answers(case, {"answers": {"q": answer}}) + off = {"A": 0.4, "B": 0.3, "C": 0.2, "D": 0.09} + with self.assertRaises(ValueError): + bench.read_answers( + case, {"answers": {"q": {"choice": "A", "probabilities": off}}} + ) + def test_duplicate_ids(self): with tempfile.TemporaryDirectory() as directory: path = Path(directory) / "requests.jsonl"