Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 31 additions & 7 deletions src/eval/structural_m3.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
from __future__ import annotations

import statistics
from fractions import Fraction
from collections import Counter
from dataclasses import dataclass, field
from datetime import datetime
Expand All @@ -41,7 +42,9 @@

# Pre-registered on #209 (M3 protocol), from the otel-fresh results; fixed before capture.
TRUTH_RETAINED_MIN = 0.50 # cause-has-telemetry truth retained (M2a's trace-only ceiling)
CANDIDATE_FRACTION_MAX = 0.33 # median |localization| / |services| over GENERATED positives (empty = 0); literal
# median |localization| / |services| over GENERATED positives (empty = 0), compared in exact rational
# arithmetic. Amended before capture from the rounded "0.33" (#209): otel-fresh M1 is exactly 1/3.
CANDIDATE_FRACTION_MAX = Fraction(1, 3)
HEALTHY_ABSTENTION_MIN = 0.58 # healthy windows with NO HYPOTHESIS GENERATED (7/12 on otel-fresh);
# not the ungated healthy_no_claim, which also counts NO_COMPATIBLE
MIN_HEALTHY_NEGATIVES = 12
Expand Down Expand Up @@ -170,13 +173,21 @@ def _rate(k: int, n: int) -> Optional[float]:
return k / n if n else None


def median_candidate_fraction(cases: list[CaseEval], arm: str) -> Optional[Fraction]:
"""Exact median of |localization| / |services| over generated positives — a ratio of integers, so the
``<= 1/3`` gate is never decided by float rounding."""
fractions = [Fraction(len(c.arms[arm].localization), max(c.n_services, 1))
for c in cases if c.positive and c.arms[arm].generated]
return statistics.median(fractions) if fractions else None


def summarize_arm(cases: list[CaseEval], arm: str) -> dict:
"""The pre-registered per-arm metrics."""
pos = [c for c in cases if c.positive]
neg = [c for c in cases if not c.positive]
pos_tel = [c for c in pos if c.cause_has_telemetry]
generated = [c for c in pos if c.arms[arm].generated]
fractions = [len(c.arms[arm].localization) / max(c.n_services, 1) for c in generated]
fraction = median_candidate_fraction(cases, arm)
identified = [c for c in cases if c.arms[arm].outcome == "identified"]
ident_correct = [c for c in identified if c.positive and c.arms[arm].localization == (c.cause,)]
return {
Expand All @@ -187,7 +198,8 @@ def summarize_arm(cases: list[CaseEval], arm: str) -> dict:
"truth_retained_cause_has_telemetry_rate": _rate(sum(c.retained(arm) for c in pos_tel), len(pos_tel)),
"median_candidates": statistics.median([len(c.arms[arm].localization) for c in generated])
if generated else None,
"median_candidate_fraction": statistics.median(fractions) if fractions else None,
"median_candidate_fraction": None if fraction is None else float(fraction),
"median_candidate_fraction_exact": _ratio(fraction),
"negatives": len(neg),
"healthy_abstained": [sum(not c.arms[arm].generated for c in neg), len(neg)],
"healthy_abstention_rate": _rate(sum(not c.arms[arm].generated for c in neg), len(neg)),
Expand Down Expand Up @@ -226,11 +238,12 @@ def decide(cases: list[CaseEval]) -> dict:
"""Apply the pre-registered decision rules to the frozen model (the full arm)."""
full = summarize_arm(cases, FULL_ARM)
retained = full["truth_retained_cause_has_telemetry_rate"]
fraction = full["median_candidate_fraction"]
fraction = median_candidate_fraction(cases, FULL_ARM) # exact, for the gate
abstention = full["healthy_abstention_rate"]
criteria = {
"truth_retained_cause_has_telemetry": [retained, TRUTH_RETAINED_MIN, ">="],
"median_candidate_fraction": [fraction, CANDIDATE_FRACTION_MAX, "<="],
# the operands the gate compares: the exact median ratio against the exact bound
"median_candidate_fraction": [_ratio(fraction), _ratio(CANDIDATE_FRACTION_MAX), "<="],
"healthy_abstention": [abstention, HEALTHY_ABSTENTION_MIN, ">="],
}
if retained is None or abstention is None:
Expand Down Expand Up @@ -301,6 +314,16 @@ def _frac(pair: list) -> str:
return f"{k}/{n}" + (f" ({k / n:.0%})" if n else "")


def _ratio(v: Optional[Fraction]) -> Optional[str]:
"""An exact ratio as ``"n/d"`` — the form the gate compares and the post prints."""
return None if v is None else f"{v.numerator}/{v.denominator}"


def _exact(v: Optional[str]) -> str:
"""Print an exact ``"n/d"`` ratio with its decimal beside it, never the rounded decimal alone."""
return "n/a" if v is None else f"{v} (≈{float(Fraction(v)):.4f})"


def _num(v: Optional[float]) -> str:
return "n/a" if v is None else f"{v:.2f}"

Expand All @@ -312,7 +335,8 @@ def render_markdown(report: dict, *, provenance: dict) -> str:
v = report["verdicts"]
lines += ["", f"**Generalization:** `{v['generalization']}` · **M2b:** `{v['m2b']}`", ""]
for name, (value, bound, op) in v["generalization_criteria"].items():
lines.append(f"- {name}: {_num(value)} (pre-registered {op} {bound})")
shown = _exact(value) if isinstance(value, str) else _num(value)
lines.append(f"- {name}: {shown} (pre-registered {op} {bound})")
if v["protocol_deviations"]:
lines += ["", "**Protocol deviations:** " + "; ".join(v["protocol_deviations"])]
lines += ["", "| metric | " + " | ".join(ARMS) + " |", "|---|" + "---|" * len(ARMS)]
Expand All @@ -321,7 +345,7 @@ def render_markdown(report: dict, *, provenance: dict) -> str:
("truth retained (all)", lambda a: _frac(a["truth_retained_all"])),
("truth retained (cause has telemetry)", lambda a: _frac(a["truth_retained_cause_has_telemetry"])),
("median candidates", lambda a: str(a["median_candidates"])),
("median candidate fraction", lambda a: _num(a["median_candidate_fraction"])),
("median candidate fraction", lambda a: _exact(a["median_candidate_fraction_exact"])),
("healthy abstention (no hypothesis)", lambda a: _frac(a["healthy_abstained"])),
("healthy no localization claim", lambda a: _frac(a["healthy_no_claim"])),
("IDENTIFIED precision", lambda a: f"{a['identified_precision'][0]}/{a['identified_precision'][1]}"),
Expand Down
49 changes: 44 additions & 5 deletions tests/unit/test_structural_m3.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
The arms are nested views of one set of inputs (M1 ⊂ M2a ⊂ M2a+M2b); each family must recover exactly
the fault it exists for, and the verdicts must apply the thresholds frozen in the M3 protocol."""
from datetime import datetime, timedelta, timezone
from fractions import Fraction

import pytest

Expand Down Expand Up @@ -200,11 +201,49 @@ def test_candidate_fraction_above_bound_fails(self):
n_loc = int(CANDIDATE_FRACTION_MAX * 10) + 1 # 4/10 > 0.33
assert decide(_corpus(n_loc=n_loc))["generalization"] == "does_not_generalize"

def test_bound_is_the_frozen_literal_so_exactly_one_third_fails(self):
# otel-fresh M1's median fraction is exactly 1/3; the frozen text is "<= 0.33", applied literally
assert decide([
_case("otel_a", "a", n_loc=1, n_services=3), _case("otel_b", "b", n_loc=1, n_services=3),
*[_healthy(i) for i in range(MIN_HEALTHY_NEGATIVES)]])["generalization"] == "does_not_generalize"
@staticmethod
def _selectivity(*loc_over_services):
pos = [_case(f"otel_p{i}", f"p{i}", n_loc=k, n_services=n) for i, (k, n) in enumerate(loc_over_services)]
return decide(pos + [_healthy(i) for i in range(MIN_HEALTHY_NEGATIVES)])

def test_exactly_one_third_passes(self):
# the amended bound (#209): <= 1/3, exact — otel-fresh M1's value is exactly 1/3
v = self._selectivity((1, 3), (1, 3))
assert v["generalization"] == "generalizes"
assert v["generalization_criteria"]["median_candidate_fraction"][1] == "1/3"

def test_one_third_as_an_even_median_is_exact(self):
# median of 1/4 and 5/12 is exactly 1/3; a float average could round either way
assert self._selectivity((1, 4), (5, 12))["generalization"] == "generalizes"

def test_just_above_one_third_fails_and_just_below_passes(self):
assert self._selectivity((334, 1000), (334, 1000))["generalization"] == "does_not_generalize"
assert self._selectivity((331, 1000), (331, 1000))["generalization"] == "generalizes"

def test_criterion_records_and_posts_the_exact_operands(self):
# 1/3 and 331/1000 pass, 167/500 fails; all three round to 0.33, so the post must show the ratio
posts = {}
for k, n in ((1, 3), (331, 1000), (334, 1000)):
v = self._selectivity((k, n), (k, n))
value, bound, op = v["generalization_criteria"]["median_candidate_fraction"]
assert (value, bound, op) == (str(Fraction(k, n)), "1/3", "<=")
md = render_markdown({**build_m3_report([]), "verdicts": v}, provenance={})
posts[value] = next(line for line in md.splitlines() if line.startswith("- median_candidate_fraction"))
assert len(set(posts.values())) == 3
assert posts["1/3"].startswith("- median_candidate_fraction: 1/3 (≈0.3333)")
assert posts["167/500"].startswith("- median_candidate_fraction: 167/500 (≈0.3340)")

def test_selectivity_row_prints_the_exact_ratio(self):
cases = [_case("otel_a", "a", n_loc=1, n_services=4), _case("otel_b", "b", n_loc=5, n_services=12),
*[_healthy(i) for i in range(MIN_HEALTHY_NEGATIVES)]]
md = render_markdown(build_m3_report(cases), provenance={})
row = next(line for line in md.splitlines() if line.startswith("| median candidate fraction"))
assert row.count("1/3 (≈0.3333)") == len(ARMS)

def test_exact_median_is_reported(self):
arm = summarize_arm([_case("otel_a", "a", n_loc=1, n_services=4), _case("otel_b", "b", n_loc=5,
n_services=12)], FULL_ARM)
assert arm["median_candidate_fraction_exact"] == "1/3"

def test_healthy_abstention_below_bound_fails(self):
generated = MIN_HEALTHY_NEGATIVES - int(HEALTHY_ABSTENTION_MIN * MIN_HEALTHY_NEGATIVES) + 1
Expand Down
Loading