"""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)}