Files
Jyotisha/scripts/research/varga_resolution_m1.py
T
jesse-uxandClaude Code bedb4d7bd1 research(rectification): archive partial varga-resolution study (BUG-1105)
Archive research scripts, regression tests, M1 results and safe M0 smoke. Keep the incomplete study and failing quick gate explicit. Exclude full M0 JSON, raw logs and unrelated oracle newline changes.

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-30 08:49:54 +08:00

222 lines
9.9 KiB
Python

#!/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())