from __future__ import annotations import math from collections import defaultdict import torch def ece(conf: list[float], correct: list[int], bins: int = 15) -> float | None: if not conf: return None score = 0.0 total = len(conf) for b in range(bins): lo = b / bins hi = (b + 1) / bins ids = [i for i, x in enumerate(conf) if x >= lo and (x < hi or (b == bins - 1 and x <= hi))] if not ids: continue acc = sum(correct[i] for i in ids) / len(ids) cf = sum(conf[i] for i in ids) / len(ids) score += len(ids) / total * abs(acc - cf) return score def option_bucket(n: int) -> str: if n <= 2: return "2" if n <= 5: return "3-5" if n <= 10: return "6-10" return ">10" def calibration_group(row: dict) -> str: """Architecture-level calibration group, intentionally not dataset-specific.""" kind = str(row.get("kind", "choice")) if kind == "choice": return f"choice:{option_bucket(int(row.get('option_count', len(row.get('target', [])))))}" return f"{kind}:*" def temperature_for_row(calibration, row: dict) -> float: if isinstance(calibration, (int, float)): return max(0.05, min(20.0, float(calibration))) if not isinstance(calibration, dict): return 1.0 if "temperature" in calibration and "groups" not in calibration and "global_temperature" not in calibration: return max(0.05, min(20.0, float(calibration.get("temperature", 1.0)))) global_t = float(calibration.get("global_temperature", calibration.get("temperature", 1.0))) group = calibration_group(row) groups = calibration.get("groups") or {} entry = groups.get(group) if isinstance(entry, dict): t = float(entry.get("temperature", global_t)) elif entry is not None: t = float(entry) else: t = global_t return max(0.05, min(20.0, t)) def _score_subset(rows: list[dict], calibration=1.0) -> dict: if not rows: return {"n": 0, "accuracy": None, "nll": None, "brier": None, "ece": None, "mean_confidence": None} conf: list[float] = [] correct: list[int] = [] nll: list[float] = [] brier: list[float] = [] for r in rows: T = temperature_for_row(calibration, r) logits = torch.tensor(r["logits"], dtype=torch.float32) / T target = torch.tensor(r["target"], dtype=torch.float32) probs = torch.softmax(logits, dim=-1) pred = int(probs.argmax()) gold = int(target.argmax()) cf = float(probs[pred]) conf.append(cf) correct.append(int(pred == gold)) nll.append(float(-(target * torch.log(probs.clamp_min(1e-9))).sum())) brier.append(float(((probs - target) ** 2).sum())) return { "n": len(rows), "accuracy": sum(correct) / len(correct), "nll": sum(nll) / len(nll), "brier": sum(brier) / len(brier), "ece": ece(conf, correct), "mean_confidence": sum(conf) / len(conf), } def score_rows(rows: list[dict], calibration=1.0, grouped: bool = True) -> dict: out = _score_subset(rows, calibration) if isinstance(calibration, (int, float)): out["temperature"] = float(calibration) elif isinstance(calibration, dict): out["calibration"] = calibration if not grouped: return out group_specs = { "by_dataset": lambda r: r.get("dataset", "unknown"), "by_kind": lambda r: r.get("kind", "choice"), "by_transform": lambda r: r.get("transform", "unknown"), "by_option_bucket": lambda r: option_bucket(int(r.get("option_count", len(r.get("target", []))))), "by_truncation": lambda r: "truncated" if r.get("truncated", False) else "not_truncated", "by_dataset_truncation": lambda r: f"{r.get('dataset','unknown')}|{'truncated' if r.get('truncated',False) else 'not_truncated'}", "by_kind_truncation": lambda r: f"{r.get('kind','choice')}|{'truncated' if r.get('truncated',False) else 'not_truncated'}", } for out_key, key_fn in group_specs.items(): groups: dict[str, list[dict]] = defaultdict(list) for r in rows: groups[str(key_fn(r))].append(r) out[out_key] = {k: _score_subset(v, calibration) for k, v in sorted(groups.items())} # Fixed-label NLI confusion is especially useful in v1.3 because the # production options are shuffled while the auxiliary semantic labels are # always entailment/neutral/contradiction. labels = ["entailment", "neutral", "contradiction"] matrix = {g: {p: 0 for p in labels} for g in labels} n_nli = 0 for r in rows: gold_idx = int(r.get("aux_nli_label", -1)) ids = list(r.get("option_ids") or []) if gold_idx < 0 or gold_idx >= 3 or not ids: continue T = temperature_for_row(calibration, r) probs = torch.softmax(torch.tensor(r["logits"], dtype=torch.float32) / T, dim=-1) pred_i = int(probs.argmax()) pred_id = str(ids[pred_i]) if pred_i < len(ids) else "" if pred_id not in labels: continue matrix[labels[gold_idx]][pred_id] += 1 n_nli += 1 if n_nli: out["nli_confusion"] = {"n": n_nli, "labels": labels, "matrix": matrix} return out def fit_temperature(rows: list[dict], max_iter: int = 50) -> float: if not rows: return 1.0 log_t = torch.tensor(0.0, requires_grad=True) opt = torch.optim.LBFGS([log_t], lr=0.2, max_iter=max_iter, line_search_fn="strong_wolfe") def closure(): opt.zero_grad() T = log_t.exp().clamp(0.05, 20.0) total = torch.tensor(0.0) for r in rows: logits = torch.tensor(r["logits"], dtype=torch.float32) target = torch.tensor(r["target"], dtype=torch.float32) total = total - (target * torch.log_softmax(logits / T, dim=-1)).sum() loss = total / max(1, len(rows)) loss.backward() return loss opt.step(closure) return float(log_t.exp().clamp(0.05, 20.0).detach()) def fit_calibration_policy( rows: list[dict], min_group: int = 40, prior_strength: float = 100.0, max_iter: int = 50, ) -> dict: """Fit shrinkage temperatures by primitive/option-count. Dataset-specific temperatures are intentionally avoided. Small groups shrink in log-temperature space toward the global fit to reduce calibration overfit. """ global_t = fit_temperature(rows, max_iter=max_iter) grouped: dict[str, list[dict]] = defaultdict(list) for row in rows: grouped[calibration_group(row)].append(row) groups: dict[str, dict] = {} for key, subset in sorted(grouped.items()): n = len(subset) if n < min_group: continue raw_t = fit_temperature(subset, max_iter=max_iter) w = n / (n + float(prior_strength)) shrunk = math.exp(w * math.log(raw_t) + (1.0 - w) * math.log(global_t)) groups[key] = { "temperature": float(max(0.05, min(20.0, shrunk))), "raw_temperature": float(raw_t), "n": n, "shrinkage_weight": w, } return { "version": 2, "type": "primitive_option_temperature", "global_temperature": float(global_t), "groups": groups, "min_group": int(min_group), "prior_strength": float(prior_strength), "n": len(rows), } def qualification_value(metrics: dict, metric: str) -> float: """Return a lower-is-better scalar for checkpoint selection.""" if metric in {"calibrated_nll", "policy_nll"}: return float(metrics["calibrated"]["nll"]) if metric in {"macro_dataset_calibrated_nll", "macro_dataset_policy_nll"}: vals = [float(x["nll"]) for x in metrics["calibrated"].get("by_dataset", {}).values() if x.get("nll") is not None] if not vals: return float(metrics["calibrated"]["nll"]) return sum(vals) / len(vals) if metric in {"macro_kind_calibrated_nll", "macro_kind_policy_nll"}: vals = [float(x["nll"]) for x in metrics["calibrated"].get("by_kind", {}).values() if x.get("nll") is not None] if not vals: return float(metrics["calibrated"]["nll"]) return sum(vals) / len(vals) if metric in {"calibrated_brier", "policy_brier"}: return float(metrics["calibrated"]["brier"]) if metric == "raw_nll": return float(metrics["raw"]["nll"]) if metric == "raw_brier": return float(metrics["raw"]["brier"]) raise ValueError(f"Unknown selection metric: {metric}")