Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017eEAG8HD3mm8gsKXgk8uU8
616 lines
35 KiB
Python
616 lines
35 KiB
Python
#!/usr/bin/env python3
|
||
"""Offline research runner for TASK-rectification-scoring-research-20260929.
|
||
|
||
R-0 research scorer reconciled per minute against the production scorer, then
|
||
per-(candidate, event) rule features recorded (eleven production rules,
|
||
auxiliary rules, and unscored extras: KP cusp sub/star lords, Pranapada /
|
||
Hora / Ghati / Bhava lagna, D60).
|
||
R-A likelihood-ratio weights (Laplace α = 1) per radius, leave-one-CASE-out
|
||
and full-fit; variants A1 (production features, reward only), A2 (+extras,
|
||
reward only), A3 (+extras, naive-Bayes present/absent, negative allowed);
|
||
six-probe replay, engine top1, calibration curve, robustness (±7-day and
|
||
±1–3-month date shifts, 1–2 flipped answers).
|
||
R-B absence evidence: gap years inside a domain's spoken span answered "no"
|
||
at 0.5 / 1.0 / 2.0 points (a card is 2.0), plus "user omitted 1–2 events"
|
||
robustness.
|
||
R-C precision follow-up: with all day events degraded to year ("spoken"
|
||
form), how many events would split candidates if the month were known,
|
||
and the gain from asking 1 / 2 of them (as a prior recompute and as an
|
||
answered probe).
|
||
|
||
Every verdict field is written as ``pending_v5``: the task brief allows a
|
||
verdict only on the v5 open set. v4 numbers are debugging references.
|
||
|
||
Run:
|
||
PYTHONHASHSEED=0 python3 scripts/research/scoring_research.py --limit 2 --radii 10 --no-write
|
||
PYTHONHASHSEED=0 python3 scripts/research/scoring_research.py
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import gc
|
||
import json
|
||
import pickle
|
||
import sys
|
||
import time as clock_module
|
||
import traceback
|
||
from collections import defaultdict
|
||
from pathlib import Path
|
||
from random import Random
|
||
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.active_rectification_event_engine import AYANAMSA, NODE_MODE, compute_candidate_static_contexts # noqa: E402
|
||
from scripts.rectification.scoring_service import build_event_contribution_matrix, score_from_matrix, scoreable_request # noqa: E402
|
||
from scripts.research.minute_resolution_sweep import MINUTE_STEP, scoring_request_for # noqa: E402
|
||
from scripts.research.offline_research_20260926_lib import exhaustive_flip_sets, flipped_answers, sample_flip_sets # noqa: E402
|
||
from scripts.research.precision_gate_lib import gate_verdict, jitter_day_events # noqa: E402
|
||
from scripts.research.probe_supply_after_six import ASK_COUNT, optimal_answer # noqa: E402
|
||
from scripts.active_rectification_events import precision_weight # noqa: E402
|
||
import scripts.research.scoring_research_lib as lib # noqa: E402
|
||
|
||
HOLDOUT_V4 = ROOT / "references" / "real_case_calibration" / "minute_rectification_holdout_v4.json"
|
||
OUT_DIR = ROOT / "docs" / "research"
|
||
RADII = (10, 30, 60)
|
||
PRIORS = ("raw", "percent")
|
||
ABSENCE_POINTS = (0.5, 1.0, 2.0) # score points per "no"; a card is 2.0
|
||
CARD_POINTS = 2.0
|
||
OMIT_REPEATS = 5
|
||
SHIFT_REPEATS = 3
|
||
FLIP_REPEATS = 10
|
||
SEED = 20260929
|
||
ABSENCE_SOURCE = "absence_research"
|
||
FOLLOWUP_SOURCE = "precision_followup_research"
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# per-case data
|
||
# ---------------------------------------------------------------------------
|
||
|
||
class CaseData:
|
||
def __init__(self, case: dict[str, Any], radius: int, cache_dir: Path | None) -> None:
|
||
self.case = case
|
||
self.case_id = str(case["case_id"])
|
||
self.radius = radius
|
||
self.true_time = str(case["birth"]["time"])[:5]
|
||
self.request = scoring_request_for({**case, "candidate_radius_minutes": radius}, radius)
|
||
self.contexts = self._contexts(cache_dir)
|
||
self.times = [ctx["candidate_at"].strftime("%H:%M") for ctx in self.contexts]
|
||
self.truth_times = lib.truth_cluster_times(self.contexts, self.true_time)
|
||
self.d60 = lib.d60_charts(self.contexts)
|
||
# production
|
||
prod_built = build_event_contribution_matrix(self.request, static_contexts=self.contexts)
|
||
self.prod_rows = score_from_matrix(self.request, prod_built)
|
||
self.prod_built = prod_built
|
||
self.probes = lib.production_probes(self.request, prod_built, self.times, self.true_time)
|
||
self.prod_public = lib.public_for(self.prod_rows, self.contexts)
|
||
# research recording pass (must reconcile to production)
|
||
self.store = lib.FeatureStore()
|
||
built = build_event_contribution_matrix(
|
||
self.request, row_provider=lib.make_recording_provider(self.contexts, self.store, self.d60),
|
||
static_contexts=self.contexts,
|
||
)
|
||
rows = score_from_matrix(self.request, built)
|
||
self.reconciliation = lib.reconcile_rows(self.prod_rows, rows)
|
||
self.training_ids = lib.training_event_ids(self.request)
|
||
self.representatives = lib.representatives_for(self.contexts)
|
||
self.events_by_id = {str(e["id"]): e for e in self.request["events"]}
|
||
|
||
def _contexts(self, cache_dir: Path | None) -> list[dict[str, Any]]:
|
||
if cache_dir is not None:
|
||
path = cache_dir / f"{self.case_id}_{self.radius}.pkl"
|
||
if path.exists():
|
||
with path.open("rb") as handle:
|
||
return pickle.load(handle)
|
||
contexts = compute_candidate_static_contexts(self.request)
|
||
if cache_dir is not None:
|
||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||
with (cache_dir / f"{self.case_id}_{self.radius}.pkl").open("wb") as handle:
|
||
pickle.dump(contexts, handle)
|
||
return contexts
|
||
|
||
# features -------------------------------------------------------------
|
||
def examples(self, *, include_extras: bool) -> list[tuple[dict[str, float], bool]]:
|
||
out: list[tuple[dict[str, float], bool]] = []
|
||
truth = set(self.truth_times)
|
||
for matrix in self.feature_matrices(include_extras=include_extras).values():
|
||
for time in self.times:
|
||
out.append((matrix.get(time, {}), time in truth))
|
||
return out
|
||
|
||
def feature_matrices(self, *, include_extras: bool) -> dict[str, dict[str, dict[str, float]]]:
|
||
cache = self.__dict__.setdefault("_fm_cache", {})
|
||
if include_extras not in cache:
|
||
cache[include_extras] = {
|
||
event_id: lib.event_feature_matrix(self.store, event_id, self.times, include_extras=include_extras)
|
||
for event_id in self.training_ids
|
||
}
|
||
return cache[include_extras]
|
||
|
||
def proxy_rows(self, table: lib.LRTable, variant: str) -> list[dict[str, Any]]:
|
||
"""Linear proxy of `weighted_rows` for fitting scale / temperature (no matrix build)."""
|
||
include_extras, allow_negative = lib.VARIANT_SPEC[variant]
|
||
universe = [n for n in sorted(table.present) if (include_extras or not lib.is_extra_feature(n))]
|
||
matrices = self.feature_matrices(include_extras=include_extras)
|
||
rows = []
|
||
for time in self.times:
|
||
total = 0.0
|
||
for event_id, matrix in matrices.items():
|
||
weight = precision_weight(str(self.events_by_id[event_id]["precision"]))
|
||
total += weight * lib.weighted_points(matrix.get(time, {}), table, allow_negative=allow_negative, universe=universe)
|
||
rows.append({"time": time, "score": round(total, 6)})
|
||
return rows
|
||
|
||
# scoring with a weight table ------------------------------------------
|
||
def weighted_rows(
|
||
self, table: lib.LRTable, variant: str, scale: float,
|
||
store: lib.FeatureStore | None = None, request: dict[str, Any] | None = None,
|
||
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
|
||
request = request or self.request
|
||
built = build_event_contribution_matrix(
|
||
request,
|
||
row_provider=lib.make_weighted_provider(store or self.store, table, variant=variant, scale=scale),
|
||
static_contexts=self.contexts,
|
||
)
|
||
return score_from_matrix(request, built), built
|
||
|
||
def rescored_store(self, events: Sequence[dict[str, Any]]) -> tuple[lib.FeatureStore, dict[str, Any], list[dict[str, Any]]]:
|
||
"""Re-record features for an altered event list (date shifts): (store, request, production rows)."""
|
||
request = scoring_request_for({**self.case, "events": list(events), "candidate_radius_minutes": self.radius}, self.radius)
|
||
store = lib.FeatureStore()
|
||
built = build_event_contribution_matrix(
|
||
request, row_provider=lib.make_recording_provider(self.contexts, store, self.d60), static_contexts=self.contexts,
|
||
)
|
||
return store, request, score_from_matrix(request, built)
|
||
|
||
def production_rows_for(self, events: Sequence[dict[str, Any]]) -> list[dict[str, Any]]:
|
||
request = scoring_request_for({**self.case, "events": list(events), "candidate_radius_minutes": self.radius}, self.radius)
|
||
built = build_event_contribution_matrix(request, static_contexts=self.contexts)
|
||
return score_from_matrix(request, built)
|
||
|
||
|
||
def evaluate(data: CaseData, rows: Sequence[dict[str, Any]], *, softmax_prior: bool, pre_answers=(), answers=None, temperature: float = 1.0, probes=None) -> dict[str, dict[str, Any]]:
|
||
public = lib.public_for(rows, data.contexts)
|
||
out: dict[str, dict[str, Any]] = {}
|
||
for mode in PRIORS:
|
||
prior = lib.prior_for(public, mode=mode, softmax=softmax_prior, temperature=temperature)
|
||
result = lib.replay(
|
||
public=public, prior=prior, probes=data.probes if probes is None else probes, true_time=data.true_time,
|
||
window_times=data.times, pre_answers=pre_answers, answers=answers,
|
||
)
|
||
result["engine_top1"] = lib.engine_top1(public, data.true_time)
|
||
out[mode] = result
|
||
return out
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# R-A
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def fit_prior(training: Sequence[CaseData], table: lib.LRTable, variant: str) -> tuple[float, float]:
|
||
"""(scale, temperature) from training cases only.
|
||
|
||
scale: research within-case score range matched to the production median
|
||
range (so the raw-prior replay's ±2 / 8-lead constants mean the same);
|
||
temperature: single Platt temperature for the percent (softmax) prior.
|
||
Training-case scores use the linear proxy (`CaseData.proxy_rows`: feature
|
||
weights × precision weight, no kind factor / transition proximity) so a
|
||
fold costs O(n) dictionary sums instead of n contribution-matrix builds.
|
||
"""
|
||
prod = [lib.within_case_range(d.prod_rows) for d in training]
|
||
rows_by_case = {d.case_id: d.proxy_rows(table, variant) for d in training}
|
||
research = [lib.within_case_range(rows) for rows in rows_by_case.values()]
|
||
p_med, r_med = lib.median(prod), lib.median(research)
|
||
scale = round(p_med / r_med, 6) if p_med and r_med else 1.0
|
||
publics = [
|
||
(lib.public_for([{**r, "score": float(r["score"]) * scale} for r in rows_by_case[d.case_id]], d.contexts), d.truth_times)
|
||
for d in training
|
||
]
|
||
return scale, lib.fit_temperature(publics)
|
||
|
||
|
||
def run_ra(by_radius: dict[int, list[CaseData]], *, fast: bool, probes_per_variant: bool = False) -> dict[str, Any]:
|
||
out: dict[str, Any] = {"variants": list(lib.VARIANTS), "alpha": lib.LAPLACE_ALPHA, "probes_per_variant": probes_per_variant, "radii": {}}
|
||
for radius, cases in sorted(by_radius.items()):
|
||
block: dict[str, Any] = {"n": len(cases), "baseline": {}, "full_fit": {}, "loo": {}, "robustness": {}, "calibration": {}, "lr_table": {}, "noise_features": {}}
|
||
baseline_rows = []
|
||
for data in cases:
|
||
row = evaluate(data, data.prod_rows, softmax_prior=False)
|
||
baseline_rows.append({"case_id": data.case_id, **{f"{m}": row[m] for m in PRIORS}})
|
||
block["baseline"] = {m: lib.summarize([r[m] for r in baseline_rows]) for m in PRIORS}
|
||
block["per_case"] = {"baseline": [{"case_id": r["case_id"], **{m: {k: r[m][k] for k in ("top1", "coverage", "width", "engine_top1")} for m in PRIORS}} for r in baseline_rows]}
|
||
per_case_examples = {
|
||
include: {d.case_id: d.examples(include_extras=include) for d in cases} for include in (False, True)
|
||
}
|
||
altered_stores: dict[tuple[str, str, int], tuple[lib.FeatureStore, dict[str, Any]]] = {}
|
||
for variant in lib.VARIANTS:
|
||
include_extras, _neg = lib.VARIANT_SPEC[variant]
|
||
examples = per_case_examples[include_extras]
|
||
# full fit -------------------------------------------------------
|
||
full_table = lib.estimate_lr([ex for rows in examples.values() for ex in rows])
|
||
full_scale, full_temperature = fit_prior(cases, full_table, variant)
|
||
full_rows = []
|
||
for data in cases:
|
||
rows, _ = data.weighted_rows(full_table, variant, full_scale)
|
||
full_rows.append(evaluate(data, rows, softmax_prior=True, temperature=full_temperature))
|
||
block["full_fit"][variant] = {
|
||
"scale": full_scale,
|
||
"temperature": full_temperature,
|
||
**{m: lib.summarize([r[m] for r in full_rows]) for m in PRIORS},
|
||
}
|
||
if variant == "A3":
|
||
block["lr_table"] = {
|
||
name: {
|
||
"log_lr_present": full_table.present[name],
|
||
"log_lr_absent": full_table.absent[name],
|
||
**full_table.support[name],
|
||
"extra": lib.is_extra_feature(name),
|
||
}
|
||
for name in sorted(full_table.present)
|
||
}
|
||
ci = lib.bootstrap_lr(examples, resamples=50 if fast else 200)
|
||
for name, interval in ci.items():
|
||
block["lr_table"].setdefault(name, {}).update(interval)
|
||
block["noise_features"] = sorted(
|
||
name for name, row in block["lr_table"].items()
|
||
if row.get("p05") is not None and row["p05"] <= 0.0 <= row["p95"]
|
||
)
|
||
# leave-one-case-out -------------------------------------------
|
||
loo_rows = []
|
||
robustness = defaultdict(list)
|
||
calibration_points: list[tuple[float, bool]] = []
|
||
temperatures: list[float] = []
|
||
for held in cases:
|
||
training = [d for d in cases if d.case_id != held.case_id]
|
||
table = lib.estimate_lr([ex for d in training for ex in examples[d.case_id]])
|
||
scale, temperature = fit_prior(training, table, variant)
|
||
temperatures.append(temperature)
|
||
rows, built_v = held.weighted_rows(table, variant, scale)
|
||
result = evaluate(held, rows, softmax_prior=True, temperature=temperature)
|
||
if probes_per_variant:
|
||
own = lib.production_probes(held.request, built_v, held.times, held.true_time)
|
||
result["own_probes"] = evaluate(held, rows, softmax_prior=True, temperature=temperature, probes=own)
|
||
result["own_probes_count"] = len(own)
|
||
result["case_id"] = held.case_id
|
||
loo_rows.append(result)
|
||
if variant == "A3":
|
||
public = lib.public_for(rows, held.contexts)
|
||
posterior = lib.percent_softmax({str(r["time"])[:5]: float(r["score"]) for r in public}, temperature)
|
||
truth = set(held.truth_times)
|
||
for rep in public:
|
||
stamp = str(rep["time"])[:5]
|
||
members = {str(t)[:5] for t in (rep.get("cluster_times") or [stamp])}
|
||
calibration_points.append((posterior.get(stamp, 0.0) / 100.0, bool(members & truth)))
|
||
# robustness: date shifts (features re-recorded with fold weights)
|
||
if not fast or variant == "A3":
|
||
events = list(held.case.get("events") or [])
|
||
shifted_sets = {
|
||
"jitter_7d": [jitter_day_events(events, 7, case_id=held.case_id)],
|
||
"shift_1_3_months": [
|
||
lib.shift_day_events_by_months(events, seed=f"{SEED}:{held.case_id}:{radius}:{k}")
|
||
for k in range(1 if fast else SHIFT_REPEATS)
|
||
],
|
||
}
|
||
for label, variants_of_events in shifted_sets.items():
|
||
for index, altered in enumerate(variants_of_events):
|
||
cache_key = (held.case_id, label, index)
|
||
if cache_key not in altered_stores:
|
||
altered_stores[cache_key] = held.rescored_store(altered)[:2]
|
||
store, request_alt = altered_stores[cache_key]
|
||
rows_alt, _ = held.weighted_rows(table, variant, scale, store=store, request=request_alt)
|
||
robustness[label].append(evaluate(held, rows_alt, softmax_prior=True, temperature=temperature))
|
||
# answer flips on the LOO rows
|
||
asked = list(held.probes)[:ASK_COUNT]
|
||
optimal = [optimal_answer(p, held.true_time) for p in asked]
|
||
for k in (1, 2):
|
||
sets = exhaustive_flip_sets(optimal, k) if k == 1 else sample_flip_sets(
|
||
optimal, k, 3 if fast else FLIP_REPEATS, case_id=held.case_id, radius=radius,
|
||
)
|
||
for picked in sets:
|
||
robustness[f"flip_{k}"].append(evaluate(
|
||
held, rows, softmax_prior=True, answers=flipped_answers(optimal, picked), temperature=temperature,
|
||
))
|
||
block["loo"][variant] = {m: lib.summarize([r[m] for r in loo_rows]) for m in PRIORS}
|
||
block["loo"][variant]["temperatures"] = sorted(set(temperatures))
|
||
block["loo"][variant]["gate_verdict"] = {
|
||
m: gate_verdict(block["baseline"][m], block["loo"][variant][m]) for m in PRIORS
|
||
}
|
||
block["loo"][variant]["hard_line_1"] = {
|
||
m: {
|
||
"coverage_count": sum(1 for r in loo_rows if r[m]["coverage"]),
|
||
"baseline_coverage_count": sum(1 for r in baseline_rows if r[m]["coverage"]),
|
||
"top1": block["loo"][variant][m]["top1"],
|
||
"baseline_top1": block["baseline"][m]["top1"],
|
||
"pass": (sum(1 for r in loo_rows if r[m]["coverage"]) >= sum(1 for r in baseline_rows if r[m]["coverage"])
|
||
and float(block["loo"][variant][m]["top1"]) + 1e-9 >= float(block["baseline"][m]["top1"])),
|
||
"note": "truth and opposite cells coincide under D1 (nothing is injected after the six probes)",
|
||
}
|
||
for m in PRIORS
|
||
}
|
||
if probes_per_variant:
|
||
block["loo"][variant]["own_probes"] = {m: lib.summarize([r["own_probes"][m] for r in loo_rows]) for m in PRIORS}
|
||
block["loo"][variant]["own_probes"]["mean_count"] = round(sum(r["own_probes_count"] for r in loo_rows) / max(len(loo_rows), 1), 2)
|
||
block["per_case"][variant] = [{"case_id": r["case_id"], **{m: {k: r[m][k] for k in ("top1", "coverage", "width", "engine_top1")} for m in PRIORS}} for r in loo_rows]
|
||
block["loo"][variant]["verdict"] = "pending_v5"
|
||
if robustness:
|
||
block["robustness"][variant] = {
|
||
label: {m: lib.summarize([r[m] for r in rows_]) for m in PRIORS}
|
||
for label, rows_ in robustness.items()
|
||
}
|
||
if variant == "A3":
|
||
block["calibration"] = lib.calibration_bins(calibration_points)
|
||
out["radii"][str(radius)] = block
|
||
return out
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# R-B
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def absence_probes(data: CaseData, events: Sequence[dict[str, Any]]) -> tuple[list[dict[str, Any]], dict[str, list[int]]]:
|
||
gaps = lib.absent_years(events)
|
||
probes: list[dict[str, Any]] = []
|
||
for domain, years in gaps.items():
|
||
for year in years:
|
||
probe = lib.split_probe(
|
||
data.representatives, birth_date=str(data.request["birth_date"]), domain=domain,
|
||
year=year, month=None, source=ABSENCE_SOURCE,
|
||
)
|
||
if probe is not None:
|
||
probes.append(probe)
|
||
return probes, gaps
|
||
|
||
|
||
def run_rb(by_radius: dict[int, list[CaseData]], *, fast: bool) -> dict[str, Any]:
|
||
out: dict[str, Any] = {"points": list(ABSENCE_POINTS), "card_points": CARD_POINTS, "omit_repeats": OMIT_REPEATS, "radii": {}}
|
||
for radius, cases in sorted(by_radius.items()):
|
||
rows_by_arm: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
||
gap_stats = []
|
||
omit_rows: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
||
for data in cases:
|
||
training = [data.events_by_id[i] for i in data.training_ids]
|
||
spoken = [{**e, "date": e.get("date_start"), "domain": e["domain"]} for e in training]
|
||
probes, gaps = absence_probes(data, spoken)
|
||
gap_stats.append({
|
||
"case_id": data.case_id, "absent_years": sum(len(v) for v in gaps.values()),
|
||
"domains_with_gaps": len(gaps), "absence_probes": len(probes),
|
||
})
|
||
base = evaluate(data, data.prod_rows, softmax_prior=False)
|
||
rows_by_arm["B0"].append(base)
|
||
for points in ABSENCE_POINTS:
|
||
pre = [(probe, "no", points / CARD_POINTS) for probe in probes]
|
||
rows_by_arm[f"absence@{points}"].append(evaluate(data, data.prod_rows, softmax_prior=False, pre_answers=pre))
|
||
# omission robustness: user did not mention 1–2 of the training events
|
||
for k in (1, 2):
|
||
for rep in range(1 if fast else OMIT_REPEATS):
|
||
kept = lib.drop_events(spoken, k, seed=f"{SEED}:{data.case_id}:{radius}:{k}:{rep}")
|
||
probes_k, _ = absence_probes(data, kept)
|
||
for points in ABSENCE_POINTS:
|
||
pre = [(probe, "no", points / CARD_POINTS) for probe in probes_k]
|
||
omit_rows[f"omit_{k}@{points}"].append(evaluate(data, data.prod_rows, softmax_prior=False, pre_answers=pre))
|
||
block: dict[str, Any] = {
|
||
"n": len(cases),
|
||
"gap_stats": {
|
||
"mean_absent_years": round(sum(g["absent_years"] for g in gap_stats) / len(gap_stats), 2),
|
||
"mean_absence_probes": round(sum(g["absence_probes"] for g in gap_stats) / len(gap_stats), 2),
|
||
"cases_with_probes": sum(1 for g in gap_stats if g["absence_probes"]),
|
||
},
|
||
"arms": {arm: {m: lib.summarize([r[m] for r in rows]) for m in PRIORS} for arm, rows in rows_by_arm.items()},
|
||
"omission": {arm: {m: lib.summarize([r[m] for r in rows]) for m in PRIORS} for arm, rows in omit_rows.items()},
|
||
}
|
||
for arm in block["arms"]:
|
||
if arm == "B0":
|
||
continue
|
||
block["arms"][arm]["gate_verdict"] = {m: gate_verdict(block["arms"]["B0"][m], block["arms"][arm][m]) for m in PRIORS}
|
||
block["arms"][arm]["hard_line_1"] = {m: {
|
||
"coverage_count": sum(1 for r in rows_by_arm[arm] if r[m]["coverage"]),
|
||
"baseline_coverage_count": sum(1 for r in rows_by_arm["B0"] if r[m]["coverage"]),
|
||
"pass": sum(1 for r in rows_by_arm[arm] if r[m]["coverage"]) >= sum(1 for r in rows_by_arm["B0"] if r[m]["coverage"])
|
||
and float(block["arms"][arm][m]["top1"]) + 1e-9 >= float(block["arms"]["B0"][m]["top1"]),
|
||
} for m in PRIORS}
|
||
block["arms"][arm]["verdict"] = "pending_v5"
|
||
out["radii"][str(radius)] = block
|
||
return out
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# R-C
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def run_rc(by_radius: dict[int, list[CaseData]], *, fast: bool) -> dict[str, Any]:
|
||
out: dict[str, Any] = {"radii": {}}
|
||
for radius, cases in sorted(by_radius.items()):
|
||
rows_by_arm: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
||
askable_counts = []
|
||
for data in cases:
|
||
originals = {str(e["id"]): e for e in data.case.get("events") or []}
|
||
raw_events = list(data.case.get("events") or [])
|
||
spoken = lib.degrade_to_year(raw_events)
|
||
spoken_rows = data.production_rows_for(spoken)
|
||
|
||
def keep(arm: str, result: dict[str, dict[str, Any]]) -> None:
|
||
rows_by_arm[arm].append({"case_id": data.case_id, **result})
|
||
|
||
keep("B0_true_precision", evaluate(data, data.prod_rows, softmax_prior=False))
|
||
keep("C0_spoken", evaluate(data, spoken_rows, softmax_prior=False))
|
||
# which spoken (year) events would split candidates if the month were known?
|
||
training = set(data.training_ids)
|
||
id_map = {str(e["id"]): str(rid) for rid, e in zip(
|
||
[str(x["id"]) for x in data.request["events"]], data.case.get("events") or [],
|
||
)}
|
||
candidates = []
|
||
for event in raw_events:
|
||
if str(event.get("precision")) not in {"day", "month"}:
|
||
continue
|
||
request_id = id_map.get(str(event["id"]))
|
||
if request_id not in training:
|
||
continue
|
||
year, month = int(str(event["date"])[:4]), int(str(event["date"])[5:7])
|
||
probe = lib.split_probe(
|
||
data.representatives, birth_date=str(data.request["birth_date"]), domain=str(event["domain"]),
|
||
year=year, month=month, source=FOLLOWUP_SOURCE,
|
||
)
|
||
if probe is None:
|
||
continue
|
||
candidates.append((float(probe.get("information_gain") or 0), str(event["id"]), probe))
|
||
candidates.sort(key=lambda item: (-item[0], item[1]))
|
||
askable_counts.append(len(candidates))
|
||
for k in (1, 2):
|
||
picked = candidates[:k]
|
||
if len(picked) < k:
|
||
continue
|
||
ids = {item[1] for item in picked}
|
||
upgraded = lib.upgrade_to_month(spoken, ids, originals)
|
||
upgraded_rows = data.production_rows_for(upgraded)
|
||
pre = [(item[2], "yes", 1.0) for item in picked]
|
||
keep(f"C{k}_prior", evaluate(data, upgraded_rows, softmax_prior=False))
|
||
keep(f"C{k}_probe", evaluate(data, spoken_rows, softmax_prior=False, pre_answers=pre))
|
||
keep(f"C{k}_both", evaluate(data, upgraded_rows, softmax_prior=False, pre_answers=pre))
|
||
by_case = {arm: {r["case_id"]: r for r in rows} for arm, rows in rows_by_arm.items()}
|
||
block: dict[str, Any] = {
|
||
"n": len(cases),
|
||
"askable_per_case": {
|
||
"mean": round(sum(askable_counts) / len(askable_counts), 2) if askable_counts else None,
|
||
"distribution": {str(k): askable_counts.count(k) for k in sorted(set(askable_counts))},
|
||
"cases_with_at_least_one": sum(1 for c in askable_counts if c >= 1),
|
||
"cases_with_at_least_two": sum(1 for c in askable_counts if c >= 2),
|
||
},
|
||
"arms": {arm: {m: lib.summarize([r[m] for r in rows]) for m in PRIORS} for arm, rows in rows_by_arm.items()},
|
||
"per_case": {arm: [{"case_id": r["case_id"], **{m: {k: r[m][k] for k in ("top1", "coverage", "width", "engine_top1")} for m in PRIORS}} for r in rows] for arm, rows in rows_by_arm.items()},
|
||
}
|
||
for arm, rows in rows_by_arm.items():
|
||
if arm.startswith("C0") or arm.startswith("B0"):
|
||
continue
|
||
ids = [r["case_id"] for r in rows]
|
||
subset_c0 = [by_case["C0_spoken"][cid] for cid in ids]
|
||
subset_b0 = [by_case["B0_true_precision"][cid] for cid in ids]
|
||
block["arms"][arm]["subset_C0"] = {m: lib.summarize([r[m] for r in subset_c0]) for m in PRIORS}
|
||
block["arms"][arm]["subset_B0"] = {m: lib.summarize([r[m] for r in subset_b0]) for m in PRIORS}
|
||
block["arms"][arm]["gate_verdict_vs_subset_C0"] = {m: gate_verdict(block["arms"][arm]["subset_C0"][m], block["arms"][arm][m]) for m in PRIORS}
|
||
block["arms"][arm]["hard_line_1_vs_subset_C0"] = {m: {
|
||
"coverage_count": sum(1 for r in rows if r[m]["coverage"]),
|
||
"subset_C0_coverage_count": sum(1 for r in subset_c0 if r[m]["coverage"]),
|
||
"pass": sum(1 for r in rows if r[m]["coverage"]) >= sum(1 for r in subset_c0 if r[m]["coverage"])
|
||
and float(block["arms"][arm][m]["top1"]) + 1e-9 >= float(block["arms"][arm]["subset_C0"][m]["top1"]),
|
||
} for m in PRIORS}
|
||
block["arms"][arm]["gate_verdict_vs_C0"] = {m: gate_verdict(block["arms"]["C0_spoken"][m], block["arms"][arm][m]) for m in PRIORS}
|
||
block["arms"][arm]["verdict"] = "pending_v5"
|
||
out["radii"][str(radius)] = block
|
||
return out
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# main
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def load_cases(path: Path) -> tuple[dict[str, Any], list[dict[str, Any]]]:
|
||
payload = json.loads(path.read_text(encoding="utf-8"))
|
||
return payload, list(payload["cases"])
|
||
|
||
|
||
def write_json(path: Path, payload: dict[str, Any]) -> None:
|
||
path.write_text(json.dumps(payload, ensure_ascii=False, indent=1, sort_keys=True) + "\n", encoding="utf-8")
|
||
|
||
|
||
def main() -> int:
|
||
parser = argparse.ArgumentParser()
|
||
parser.add_argument("--holdout", default=str(HOLDOUT_V4))
|
||
parser.add_argument("--dataset-label", default="v4")
|
||
parser.add_argument("--limit", type=int, default=0)
|
||
parser.add_argument("--radii", nargs="+", type=int, default=list(RADII))
|
||
parser.add_argument("--parts", nargs="+", default=["ra", "rb", "rc"])
|
||
parser.add_argument("--cache-dir", default="")
|
||
parser.add_argument("--fast", action="store_true", help="fewer repeats (development)")
|
||
parser.add_argument("--no-write", action="store_true")
|
||
parser.add_argument("--out-suffix", default="2026_09_29")
|
||
parser.add_argument("--probes-per-variant", action="store_true", help="also replay LOO variants with probes generated from their own matrix")
|
||
args = parser.parse_args()
|
||
started = clock_module.time()
|
||
holdout_path = Path(args.holdout)
|
||
meta, cases = load_cases(holdout_path)
|
||
if args.limit:
|
||
cases = cases[: args.limit]
|
||
cache_dir = Path(args.cache_dir) if args.cache_dir else None
|
||
errors: list[dict[str, Any]] = []
|
||
reconciliation: list[dict[str, Any]] = []
|
||
parts = {"ra": "lr", "rb": "absence", "rc": "precision_followup"}
|
||
radii_results: dict[str, dict[str, dict[str, Any]]] = {key: {} for key in parts.values()}
|
||
part_meta: dict[str, dict[str, Any]] = {}
|
||
# One radius at a time: 77 cases x 3 radii of static contexts do not fit in memory together.
|
||
for radius in args.radii:
|
||
loaded: list[CaseData] = []
|
||
for case in cases:
|
||
try:
|
||
data = CaseData(case, radius, cache_dir)
|
||
except Exception as exc: # noqa: BLE001
|
||
errors.append({"case_id": case.get("case_id"), "radius": radius, "error": f"{type(exc).__name__}: {exc}", "trace": traceback.format_exc()})
|
||
continue
|
||
reconciliation.append({"case_id": data.case_id, "radius": radius, **data.reconciliation})
|
||
loaded.append(data)
|
||
print(f"loaded radius {radius}: {len(loaded)} cases ({clock_module.time() - started:.0f}s)", flush=True)
|
||
if any(row["changed"] for row in reconciliation if row["radius"] == radius):
|
||
print(json.dumps({"UNRECONCILED": [row for row in reconciliation if row["changed"]]}, ensure_ascii=False, indent=1))
|
||
return 2
|
||
by_radius = {radius: loaded}
|
||
runners = {
|
||
"ra": lambda: run_ra(by_radius, fast=args.fast, probes_per_variant=args.probes_per_variant),
|
||
"rb": lambda: run_rb(by_radius, fast=args.fast),
|
||
"rc": lambda: run_rc(by_radius, fast=args.fast),
|
||
}
|
||
for part, key in parts.items():
|
||
if part not in args.parts:
|
||
continue
|
||
block = runners[part]()
|
||
radii_results[key].update(block.pop("radii"))
|
||
part_meta[key] = block
|
||
print(f"{part.upper()} radius {radius} done {clock_module.time() - started:.0f}s", flush=True)
|
||
del loaded, by_radius
|
||
gc.collect()
|
||
header = {
|
||
"generated_for": "TASK-rectification-scoring-research-20260929",
|
||
"dataset": args.dataset_label,
|
||
"holdout": str(holdout_path.relative_to(ROOT)) if holdout_path.is_relative_to(ROOT) else str(holdout_path),
|
||
"case_count": len(cases),
|
||
"radii": list(args.radii),
|
||
"minute_step": MINUTE_STEP,
|
||
"ayanamsa": AYANAMSA,
|
||
"node_mode": NODE_MODE,
|
||
"ask_count": ASK_COUNT,
|
||
"probe_today": lib.TODAY.isoformat(),
|
||
"probes": "production G0 probes generated once per case/radius from the production matrix; reused for every arm (own_probes = regenerated from the variant matrix)",
|
||
"open_set_not_blind": True,
|
||
"verdict": "pending_v5" if args.dataset_label != "v5" else "see_report",
|
||
"reconciliation": {"cases": len(reconciliation), "unreconciled": sum(1 for r in reconciliation if r["changed"]), "rows": reconciliation},
|
||
"errors": errors,
|
||
"seed": SEED,
|
||
"fast": bool(args.fast),
|
||
}
|
||
results: dict[str, dict[str, Any]] = {}
|
||
for part, key in parts.items():
|
||
if part in args.parts:
|
||
results[key] = {**header, "part": {"ra": "R-A", "rb": "R-B", "rc": "R-C"}[part], **part_meta.get(key, {}), "radii": radii_results[key]}
|
||
for key, payload in results.items():
|
||
if not args.no_write:
|
||
write_json(OUT_DIR / f"scoring_research_{key}_{args.out_suffix}.json", payload)
|
||
brief: dict[str, Any] = {"reconciliation_unreconciled": header["reconciliation"]["unreconciled"], "errors": len(errors), "elapsed_seconds": round(clock_module.time() - started, 1)}
|
||
for key, payload in results.items():
|
||
brief[key] = {}
|
||
for radius, block in payload["radii"].items():
|
||
if key == "lr":
|
||
brief[key][radius] = {
|
||
"baseline": {m: (block["baseline"][m]["top1"], block["baseline"][m]["coverage"], block["baseline"][m]["width_median"], block["baseline"][m]["engine_top1"]) for m in PRIORS},
|
||
**{v: {m: (block["loo"][v][m]["top1"], block["loo"][v][m]["coverage"], block["loo"][v][m]["width_median"], block["loo"][v][m]["engine_top1"]) for m in PRIORS} for v in lib.VARIANTS},
|
||
}
|
||
else:
|
||
brief[key][radius] = {arm: {m: (s[m]["top1"], s[m]["coverage"], s[m]["width_median"]) for m in PRIORS} for arm, s in block["arms"].items()}
|
||
print(json.dumps(brief, ensure_ascii=False, indent=1))
|
||
return 0 if not errors else 1
|
||
|
||
|
||
if __name__ == "__main__":
|
||
sys.exit(main())
|