#!/usr/bin/env python3 """Offline M1 segment-quality and leave-one-out analysis for BUG-1105. This deliberately reuses ``native_case`` and does not modify production rectification code. The output contains both full-sample fixed-threshold figures and case-level leave-one-out validation; no threshold is selected from the validation case itself. """ from __future__ import annotations import argparse import json import sys from pathlib import Path from typing import Any, Sequence ROOT = Path(__file__).resolve().parents[2] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from scripts.research.varga_resolution_lib import ( # noqa: E402 RADII, THRESHOLDS, VARGA_PREFIXES, native_case, segment_metrics, choose_loo_threshold, threshold_scan, ) HOLDOUT = ROOT / "references" / "real_case_calibration" / "minute_rectification_holdout_v5.json" SCHEMA = "bug-1105-varga-resolution-m1-v1" MODES = ("raw", "percent", "uniform") def load_cases() -> list[dict[str, Any]]: return list(json.loads(HOLDOUT.read_text(encoding="utf-8")).get("cases") or []) def is_lmt(case: dict[str, Any]) -> bool: return int(str(case.get("birth", {}).get("date", "9999"))[:4]) < 1900 def metrics_for_case(case: dict[str, Any], radius: int) -> dict[str, Any]: result = native_case(case, radius, do_reconcile=False) state = result["state"] rows = result["chart_rows"] true_time = result["true_time"] return { "case_id": str(case["case_id"]), "radius": radius, "lmt_before_1900": is_lmt(case), "answered_count": int((state.get("result") or {}).get("questions") or 0), "probe_count": len(result.get("probes") or []), "by_varga": { prefix: { mode: segment_metrics(state, rows, prefix, true_time, mode) for mode in MODES } for prefix in VARGA_PREFIXES }, } def aggregate_rows(items: Sequence[dict[str, Any]], prefix: str, mode: str) -> dict[str, Any]: rows = [item["by_varga"][prefix][mode] for item in items] denominator = len(rows) return { "denominator": denominator, "top_share_mean": round(sum(float(r["top_share"] or 0) for r in rows) / denominator, 8) if denominator else None, "truth_retained": sum(bool(r["truth_retained"]) for r in rows), "truth_retained_rate": round(sum(bool(r["truth_retained"]) for r in rows) / denominator, 8) if denominator else None, "truth_excluded": sum(not bool(r["truth_retained"]) for r in rows), "top_segment_correct": sum(bool(r["top_segment_correct"]) for r in rows), "top_segment_correct_rate": round(sum(bool(r["top_segment_correct"]) for r in rows) / denominator, 8) if denominator else None, "top_segment_ties": sum(bool(r.get("top_segment_tie")) for r in rows), "top_segment_tie_rate": round(sum(bool(r.get("top_segment_tie")) for r in rows) / denominator, 8) if denominator else None, "valid_segment_count_le_2": sum(int(r["valid_segment_count"]) <= 2 for r in rows), "valid_segment_count_le_2_rate": round(sum(int(r["valid_segment_count"]) <= 2 for r in rows) / denominator, 8) if denominator else None, "thresholds_full_fit": threshold_scan(rows), } def stratified_rows(items: Sequence[dict[str, Any]], prefix: str, mode: str) -> dict[str, Any]: """Keep the pre-1900 LMT stratum auditable without changing denominators.""" return { "all": aggregate_rows(items, prefix, mode), "lmt_before_1900": aggregate_rows( [item for item in items if item["lmt_before_1900"]], prefix, mode ), "post_1900": aggregate_rows( [item for item in items if not item["lmt_before_1900"]], prefix, mode ), } def summarize_loo_rows(rows: Sequence[dict[str, Any]]) -> dict[str, Any]: eligible = [row for row in rows if row["eligible"]] return { "eligible": len(eligible), "denominator": len(rows), "coverage": round(len(eligible) / len(rows), 8) if rows else None, "accuracy": round(sum(bool(row["top_segment_correct"]) for row in eligible) / len(eligible), 8) if eligible else None, "truth_retained": round(sum(bool(row["truth_retained"]) for row in eligible) / len(eligible), 8) if eligible else None, } def loo_fixed(items: Sequence[dict[str, Any]], prefix: str, mode: str) -> dict[str, Any]: out: dict[str, Any] = {} for threshold in THRESHOLDS: selected = [item["by_varga"][prefix][mode] for item in items if float(item["by_varga"][prefix][mode]["top_share"] or 0) >= threshold] out[str(threshold)] = { "eligible": len(selected), "denominator": len(items), "coverage": round(len(selected) / len(items), 8) if items else None, "accuracy": round(sum(bool(row["top_segment_correct"]) for row in selected) / len(selected), 8) if selected else None, "truth_retained": round(sum(bool(row["truth_retained"]) for row in selected) / len(selected), 8) if selected else None, } return out def loo_selected(items: Sequence[dict[str, Any]], prefix: str, mode: str) -> dict[str, Any]: validations: list[dict[str, Any]] = [] for index, item in enumerate(items): training = [other["by_varga"][prefix][mode] for j, other in enumerate(items) if j != index] threshold = choose_loo_threshold(training) row = item["by_varga"][prefix][mode] eligible = threshold is not None and float(row["top_share"] or 0) >= threshold validations.append({ "case_id": item["case_id"], "threshold": threshold, "eligible": bool(eligible), "top_segment_correct": bool(row["top_segment_correct"]) if eligible else None, "truth_retained": bool(row["truth_retained"]) if eligible else None, "lmt_before_1900": item["lmt_before_1900"], }) eligible = [row for row in validations if row["eligible"]] return { "eligible": len(eligible), "denominator": len(validations), "coverage": round(len(eligible) / len(validations), 8) if validations else None, "accuracy": round(sum(bool(row["top_segment_correct"]) for row in eligible) / len(eligible), 8) if eligible else None, "truth_retained": round(sum(bool(row["truth_retained"]) for row in eligible) / len(eligible), 8) if eligible else None, "selected_threshold_counts": {str(t): sum(row["threshold"] == t for row in validations) for t in THRESHOLDS}, "validation_rows": validations, } def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("--radii", default=",".join(str(value) for value in RADII)) parser.add_argument("--limit", type=int, default=0) parser.add_argument("--json-out", required=True) args = parser.parse_args() radii = tuple(int(value) for value in str(args.radii).split(",") if value.strip()) cases = load_cases() if args.limit: cases = cases[: args.limit] items: list[dict[str, Any]] = [] errors: list[dict[str, Any]] = [] for case in cases: for radius in radii: label = f"{case.get('case_id')} ±{radius}" try: items.append(metrics_for_case(case, radius)) print(label, flush=True) except Exception as exc: # noqa: BLE001 errors.append({"case_id": str(case.get("case_id")), "radius": radius, "error": f"{type(exc).__name__}: {exc}"}) print(label, "ERROR", type(exc).__name__, exc, flush=True) aggregates: list[dict[str, Any]] = [] for radius in radii: subset = [item for item in items if item["radius"] == radius] by_varga: dict[str, Any] = {} for prefix in VARGA_PREFIXES: by_varga[prefix] = {} for mode in MODES: loo = loo_selected(subset, prefix, mode) by_varga[prefix][mode] = { "full_fit": stratified_rows(subset, prefix, mode), "loo_fixed_thresholds": loo_fixed(subset, prefix, mode), "loo_selected_threshold": { **summarize_loo_rows(loo["validation_rows"]), "selected_threshold_counts": loo["selected_threshold_counts"], "validation_rows": loo["validation_rows"], "lmt_before_1900": summarize_loo_rows([row for row in loo["validation_rows"] if row["lmt_before_1900"]]), "post_1900": summarize_loo_rows([row for row in loo["validation_rows"] if not row["lmt_before_1900"]]), }, } aggregates.append({ "radius": radius, "case_count": len(subset), "lmt_case_count": sum(item["lmt_before_1900"] for item in subset), "by_varga": by_varga, }) payload = { "schema": SCHEMA, "metadata": { "holdout": str(HOLDOUT.relative_to(ROOT)).replace("\\", "/"), "case_count_requested": len(cases), "case_count_completed": len({item["case_id"] for item in items}), "radii": list(radii), "vargas": list(VARGA_PREFIXES), "modes": list(MODES), "thresholds": list(THRESHOLDS), "loo": "one case held out; threshold selected only from the other cases; validation case never used for selection", "lmt_stratum": "birth year < 1900, reported separately", "production_code_modified": False, "deterministic_json": True, }, "aggregates": aggregates, "items": items, "errors": errors, } out = Path(args.json_out) out.parent.mkdir(parents=True, exist_ok=True) out.write_text(json.dumps(payload, ensure_ascii=False, indent=2, sort_keys=True) + "\n", encoding="utf-8") return 0 if not errors and len(items) == len(cases) * len(radii) else 1 if __name__ == "__main__": raise SystemExit(main())