Ines-1 / training /code /metrics.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit b7f5644)
61b6fb9 verified
Raw History Blame Contribute Delete
8.17 kB
"""Accuracy and calibration for typed decisions, from per-item outputs (eval_items.py).
Definitions (pinned by reproducing the typed-decisions card's reference rows, see `reference_rows`):
accuracy argmax == the gold label (ties broken by option order)
nll -log p(gold label) (hard target)
kl KL(gold || p) = sum_k g_k log(g_k / p_k), natural log (soft target, teacher mean)
ce_soft -sum_k g_k log p_k (soft target)
brier sum_k (p_k - g_k)^2 against the gold DISTRIBUTION, mean over questions
ece 10 equal-width bins of top-1 confidence; |accuracy - confidence| weighted by bin size
Variants are kept so a mismatch with a published definition is visible, not hidden.
"""
from __future__ import annotations
import collections
import json
import math
import random
EPS = 1e-12
def load(path):
return [json.loads(x) for x in open(path, encoding="utf-8") if x.strip()]
def soft_vec(it):
s = it.get("soft")
if not s:
return None
v = [float(s.get(k, 0.0)) for k in it["keys"]]
t = sum(v)
return [x / t for x in v] if t > 0 else None
def item_scores(it):
p, keys = it["probs"], it["keys"]
gold = it["gold"][0]
gi = keys.index(gold) if gold in keys else None
top = max(range(len(p)), key=p.__getitem__)
g = soft_vec(it)
out = {"ok": int(keys[top] in it["gold"]), "conf": p[top], "nll": -math.log(max(p[gi], EPS)) if gi is not None else None}
if g:
out["kl"] = sum(gk * math.log(gk / max(pk, EPS)) for gk, pk in zip(g, p) if gk > 0)
out["ce_soft"] = -sum(gk * math.log(max(pk, EPS)) for gk, pk in zip(g, p) if gk > 0)
out["brier"] = sum((pk - gk) ** 2 for gk, pk in zip(g, p))
out["brier_mean_k"] = out["brier"] / len(p)
oh = [1.0 if i == gi else 0.0 for i in range(len(p))]
out["brier_onehot"] = sum((pk - ok) ** 2 for pk, ok in zip(p, oh))
return out
def ece(scores, bins=10):
b = collections.defaultdict(list)
for s in scores:
b[min(int(s["conf"] * bins), bins - 1)].append(s)
n = len(scores)
return sum(len(v) / n * abs(sum(x["ok"] for x in v) / len(v) - sum(x["conf"] for x in v) / len(v)) for v in b.values())
def summarize(items):
sc = [item_scores(it) for it in items]
n = len(sc)
def mean(k):
v = [s[k] for s in sc if s.get(k) is not None]
return sum(v) / len(v) if v else None
out = {"n": n, "correct": sum(s["ok"] for s in sc), "accuracy": sum(s["ok"] for s in sc) / n,
"nll": mean("nll"), "kl": mean("kl"), "ce_soft": mean("ce_soft"), "brier": mean("brier"),
"brier_mean_k": mean("brier_mean_k"), "brier_onehot": mean("brier_onehot"), "ece10": ece(sc),
"mean_conf": sum(s["conf"] for s in sc) / n}
by = collections.defaultdict(list)
for it, s in zip(items, sc):
by[it["type"]].append(s["ok"])
out["by_type"] = {k: {"correct": sum(v), "n": len(v), "accuracy": sum(v) / len(v)} for k, v in sorted(by.items())}
return out
def bootstrap_ci(items, reps=4000, seed=0):
"""95 % interval of accuracy, resampling CASES (the questions of a case are not independent)."""
cases = collections.defaultdict(list)
for it in items:
cases[it["id"]].append(item_scores(it)["ok"])
ids = list(cases)
rng = random.Random(seed)
acc = []
for _ in range(reps):
s = [cases[rng.choice(ids)] for _ in ids]
acc.append(sum(map(sum, s)) / sum(map(len, s)))
acc.sort()
return acc[int(0.025 * reps)], acc[int(0.975 * reps) - 1]
def paired(a_items, b_items):
"""McNemar on the same questions (exact two-sided binomial on the discordant pairs)."""
A = {(i["id"], i["qid"]): item_scores(i)["ok"] for i in a_items}
B = {(i["id"], i["qid"]): item_scores(i)["ok"] for i in b_items}
common = A.keys() & B.keys()
b01 = sum(1 for k in common if not A[k] and B[k])
b10 = sum(1 for k in common if A[k] and not B[k])
n = b01 + b10
k = min(b01, b10)
p = min(1.0, 2 * sum(math.comb(n, i) for i in range(k + 1)) / 2 ** n) if n else 1.0
return {"n_common": len(common), "only_b_right": b01, "only_a_right": b10, "p_mcnemar_exact": p}
# ------------------------------------------------------------------ class-balance diagnostics
def confusion(gold, pred, labels=None):
"""labels in a fixed order (default: sorted union); returns (labels, matrix[gold][pred])."""
labels = list(labels or sorted(set(gold) | set(pred)))
ix = {l: i for i, l in enumerate(labels)}
m = [[0] * len(labels) for _ in labels]
for g, p in zip(gold, pred):
m[ix[g]][ix[p]] += 1
return labels, m
def per_class(labels, m):
out = {}
for i, l in enumerate(labels):
tp = m[i][i]
fn = sum(m[i]) - tp
fp = sum(row[i] for row in m) - tp
rec = tp / (tp + fn) if tp + fn else None
prec = tp / (tp + fp) if tp + fp else 0.0
f1 = 2 * prec * rec / (prec + rec) if rec and (prec + rec) else 0.0
out[l] = {"support": tp + fn, "predicted": tp + fp, "recall": rec, "precision": prec, "f1": f1}
return out
def balanced_metrics(gold, pred, labels=None):
"""accuracy, balanced accuracy (= macro recall over classes present in gold), macro-F1 over classes
present in gold, MCC (Gorodkin's multiclass form; equals the binary MCC for two classes)."""
labels, m = confusion(gold, pred, labels)
pc = per_class(labels, m)
present = [l for l in labels if pc[l]["support"]]
n = len(gold)
acc = sum(m[i][i] for i in range(len(labels))) / n
bal = sum(pc[l]["recall"] for l in present) / len(present)
mf1 = sum(pc[l]["f1"] for l in present) / len(present)
t = [sum(r) for r in m] # true counts
p = [sum(r[j] for r in m) for j in range(len(labels))] # predicted counts
c = sum(m[i][i] for i in range(len(labels)))
num = c * n - sum(tk * pk for tk, pk in zip(t, p))
den = math.sqrt((n * n - sum(pk * pk for pk in p)) * (n * n - sum(tk * tk for tk in t)))
mcc = num / den if den else 0.0
return {"n": n, "accuracy": acc, "balanced_accuracy": bal, "macro_f1": mf1, "mcc": mcc,
"labels": labels, "confusion": m, "gold_dist": dict(zip(labels, t)), "pred_dist": dict(zip(labels, p)),
"per_class": pc}
def auroc(scores, positives):
"""Probability that a random positive scores above a random negative (ties count 1/2)."""
pos = [s for s, y in zip(scores, positives) if y]
neg = [s for s, y in zip(scores, positives) if not y]
if not pos or not neg:
return None
order = sorted(scores)
ranks = {}
i = 0
while i < len(order): # average ranks over ties
j = i
while j < len(order) and order[j] == order[i]:
j += 1
ranks[order[i]] = (i + j + 1) / 2
i = j
rsum = sum(ranks[s] for s in pos)
return (rsum - len(pos) * (len(pos) + 1) / 2) / (len(pos) * len(neg))
def average_precision(scores, positives):
"""PR-AUC as average precision (step function, no interpolation) for the positive class."""
pairs = sorted(zip(scores, positives), key=lambda x: -x[0])
tp, ap, npos = 0, 0.0, sum(1 for _, y in pairs if y)
if not npos:
return None
for k, (_, y) in enumerate(pairs, 1):
if y:
tp += 1
ap += tp / k
return ap / npos
def score_mae(items):
"""OUR definition (no public code to match): mean |E_p[level] - E_gold[level]| over `score` questions,
levels 0..K-1, gold expectation from the teacher distribution; and against the gold label."""
d_soft, d_lab = [], []
for it in items:
if it["type"] != "score":
continue
p = it["probs"]
e = sum(i * x for i, x in enumerate(p))
g = soft_vec(it)
if g:
d_soft.append(abs(e - sum(i * x for i, x in enumerate(g))))
d_lab.append(abs(e - int(it["gold"][0])))
return {"mae_vs_gold_expectation": sum(d_soft) / len(d_soft) if d_soft else None,
"mae_vs_gold_label": sum(d_lab) / len(d_lab) if d_lab else None, "n": len(d_lab)}