"""Item 20 phase 3 (offline, no GPU): threshold sweep + oracle bound. Inputs: gate_probs.json (per-item halt probabilities, phase 1) and eval_code_merge_lut.json (per-item outcomes at k=0/1/2/4, phase 2). For each threshold theta: k*(item) = first point where cumulative halt mass >= theta (pre-loop consult included, remainder to kmax), outcome looked up per item; reports overall/easy/hard accuracy and E[k]. Also: the oracle gate (best k per item) — the gating ceiling. """ import json import os from pathlib import Path OUT = Path(os.environ.get("LOOP_OUT", Path(__file__).resolve().parent.parent / "results-loop")) probs = {r["task_id"]: r for r in json.load(open(OUT / "gate_probs.json"))} lut = json.load(open(OUT / "eval_code_merge_lut.json")) # lut["ks"][k]["per_item"] -> task_id, ok ok = {} for k, v in lut["ks"].items(): for r in v["per_item"]: ok.setdefault(r["task_id"], {})[int(k)] = r["ok"] KS = sorted(int(k) for k in lut["ks"]) # e.g. [0, 1, 2, 4] KMAX = max(KS) def kstar(p, theta): """p = [p0, p1, ..., pkmax]; deploy semantics of halted_k_per_item.""" cum, keep = 0.0, 1.0 for i, pi in enumerate(p): cum += keep * pi keep *= (1 - pi) if cum >= theta: return 0 if i == 0 else min(i, KMAX) return KMAX def nearest_k(k): """Map k* to the nearest evaluated k (lut has only KS).""" return min(KS, key=lambda x: (abs(x - k), x)) def score(assign): n, hit = {}, {} ek = 0.0 for tid, k in assign.items(): lbl = probs[tid]["label"] n[lbl] = n.get(lbl, 0) + 1 hit[lbl] = hit.get(lbl, 0) + int(ok[tid][nearest_k(k)]) ek += k tot_n = sum(n.values()) tot = sum(hit.values()) / tot_n return tot, {l: hit[l] / n[l] for l in n}, ek / tot_n print(f"{'theta':>6} {'overall':>8} {'easy':>6} {'hard':>6} {'E[k]':>5}") best = [] for theta in (0.3, 0.5, 0.7, 0.8, 0.9, 0.95, 0.98, 0.99): assign = {tid: kstar(r["p"], theta) for tid, r in probs.items()} tot, by, ek = score(assign) print(f"{theta:>6} {tot:>8.3f} {by.get('easy',0):>6.3f} " f"{by.get('hard',0):>6.3f} {ek:>5.2f}") best.append({"theta": theta, "overall": tot, "by_label": by, "ek": ek}) oracle = {tid: max(KS, key=lambda k: (ok[tid][k], -k)) for tid in probs} tot, by, ek = score(oracle) print(f"{'oracle':>6} {tot:>8.3f} {by.get('easy',0):>6.3f} " f"{by.get('hard',0):>6.3f} {ek:>5.2f}") json.dump({"curve": best, "oracle": {"overall": tot, "by_label": by, "ek": ek}}, open(OUT / "gate_threshold_curve.json", "w"), indent=1) print("wrote", OUT / "gate_threshold_curve.json")