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>
222 lines
9.9 KiB
Python
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())
|