JEV / code /scripts /fit_temperature.py
cloudyu's picture
Upload JEV v0.7.0 (Qwen3.5-9B distilled from Jev 1.13)
b2f3bf4 verified
Raw
History Blame Contribute Delete
3.44 kB
#!/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()