openjev-e4b / eval /harness /src /metrics.py
bambamdevs's picture
Publish OpenJEV E4B 1.0
03223d7
Raw History Blame Contribute Delete
8.65 kB
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}")