#!/usr/bin/env python3 """Temperature calibration on the *calibration* split only (DESIGN §8.1). python3 scripts/fit_temperature.py --checkpoint checkpoints/s2_seed42/best --out checkpoints/s2_seed42/best/calibration.json """ from __future__ import annotations import argparse import json import os import sys import numpy as np import torch sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "src")) from jev_judge import calibration as calib # noqa: E402 from jev_judge.checkpointing import load_checkpoint # noqa: E402 from jev_judge.data import load_split # noqa: E402 from jev_judge.infer import run_inference # noqa: E402 from jev_judge.metrics import per_sample_table, summarize # noqa: E402 from jev_judge.model import JevJudge, masked_probs # noqa: E402 def main() -> None: ap = argparse.ArgumentParser() g = ap.add_mutually_exclusive_group(required=True) g.add_argument("--checkpoint") g.add_argument("--base") ap.add_argument("--data", default="data") ap.add_argument("--split", default="calibration") ap.add_argument("--out", required=True) ap.add_argument("--ece-threshold", type=float, default=0.02) ap.add_argument("--min-family-n", type=int, default=500) ap.add_argument("--include-d1", action="store_true", help="keep yuri_v1 exact-uniform placeholder rows (default: excluded, consistent with the training D1 policy)") args = ap.parse_args() if args.split.startswith("test") or args.split == "ood": raise SystemExit("refusing to fit temperatures on a test/ood split (DESIGN §8.1)") judge = load_checkpoint(args.checkpoint)[0] if args.checkpoint else JevJudge.from_base(args.base) df = load_split(args.data, args.split) n_all = len(df) if not args.include_d1: from jev_judge.data import d1_flags df = df.loc[~d1_flags(df, "yuri_v1")].reset_index(drop=True) print(f"calibration rows: {len(df)} (of {n_all}; D1 placeholders {'kept' if args.include_d1 else 'excluded'})") out = run_inference(judge, df, desc=args.split) kinds = df["kind"].to_numpy() fams = df["family"].to_numpy() table = calib.fit_temperatures(out["logits"], out["q"], out["mask"], kinds, fams, args.ece_threshold, args.min_family_n) # before/after summary on the calibration split itself (diagnostic only) p_raw = masked_probs(torch.as_tensor(out["logits"]), torch.as_tensor(out["mask"])).numpy() z_t = calib.apply_temperatures(torch.as_tensor(out["logits"]), torch.as_tensor(out["kind_ids"]), table, list(fams)) p_cal = masked_probs(z_t, torch.as_tensor(out["mask"])).numpy() s_raw = summarize(per_sample_table(p_raw, out["q"], out["mask"], df)) s_cal = summarize(per_sample_table(p_cal, out["q"], out["mask"], df)) table["diagnostic_calibration_split"] = {"raw": {k: s_raw[k] for k in ("kl", "ece", "mce")}, "calibrated": {k: s_cal[k] for k in ("kl", "ece", "mce")}} table["source"] = args.checkpoint or args.base table["fit_rows"] = int(len(df)) table["d1_excluded"] = not args.include_d1 os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True) calib.save(table, args.out) print(json.dumps({"per_kind": table["per_kind"], "per_kind_family": table["per_kind_family"], "diag": table["diagnostic_calibration_split"]}, indent=2)) print("->", args.out) if __name__ == "__main__": main()