Download runtime/jevlike/evaluate.py from Cem13/kodama-core: direct link, hf CLI and curl.
- Browser
- Download file 31.2 kB
-
https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/evaluate.py
- Command line
-
hf download hf://Cem13/kodama-core/runtime/jevlike/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/evaluate.py
31.2 kB
| """Evaluate a SystemOne checkpoint on Record JSONL files. | |
| python -m jevlike.evaluate --checkpoint checkpoints/jevlike-base \ | |
| --data data/test.jsonl data/heldout.jsonl --out reports/ | |
| Writes reports/<name>.md, reports/<name>.json and reports/<name>_reliability.png. | |
| Metrics: accuracy (overall / by type / by source / by file), NLL, Brier, ECE (15 equal-width | |
| bins, top-label) for calibrated and - when the model exposes them - raw (uncalibrated) | |
| probabilities; reliability diagram; escalate-head quality (AUROC for detecting wrong answers, | |
| accuracy at 50/80/100% coverage, flag statistics) against a max-probability baseline; | |
| class-prior (majority) and uniform-random baselines; single-call latency p50/p95 and batched | |
| throughput. | |
| Records that share the same state are sent as ONE predict call with several questions, which is | |
| how the model is meant to be used. | |
| Conventions | |
| ----------- | |
| * Brier = sum_k (p_k - y_k)^2 over the question's options (noul: 2 options false/true, so it is | |
| 2*(p - y)^2), averaged over questions. | |
| * noul is scored as the 2-way distribution [P(false), P(true)]; its prediction is P(true) > 0.5. | |
| * The prior baseline predicts, for each (source, type), the label-index frequencies (add-one | |
| smoothed, restricted to the question's options) - i.e. the majority class for fixed-label | |
| datasets and the majority position/boolean otherwise. It is fitted on --prior-data if given, | |
| else in-sample on the evaluated records (optimistic). | |
| * The uniform baseline reports expected values: acc = mean 1/n, NLL = mean log n, Brier = mean 1-1/n. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import datetime as _dt | |
| import inspect | |
| import json | |
| import math | |
| import random | |
| import time | |
| from collections import Counter, defaultdict | |
| from pathlib import Path | |
| from typing import Any, Optional | |
| import numpy as np | |
| from jevlike.types import Choice, Noul, Question, Record, Score, read_jsonl | |
| N_BINS = 15 | |
| COVERAGES = (0.5, 0.8, 1.0) | |
| EPS = 1e-12 | |
| # kwargs a predict_batch may accept to return uncalibrated (temperature = 1) probabilities | |
| _RAW_KWARGS = (("calibrated", False), ("calibrate", False), ("apply_calibration", False), | |
| ("use_calibration", False), ("raw", True)) | |
| _ESCALATE_PROB_KEYS = ("escalate_prob", "escalate_probability", "p_escalate", "escalate_p", "escalate_score") | |
| _P_CORRECT_KEYS = ("p_correct", "correct_prob", "act_prob", "p_act") | |
| # ---------------------------------------------------------------- record <-> API | |
| def question_to_public(q: Question): | |
| """Normalized Question -> the public type SystemOne.predict takes (round-trips via from_public).""" | |
| if q.type == "choice": | |
| return Choice(q.instructions, {o.key: o.text for o in q.options}) | |
| if q.type == "score": | |
| return Score(q.instructions, [o.text for o in q.options]) | |
| if q.type == "noul": | |
| return Noul(q.instructions) | |
| raise ValueError(q.type) | |
| def answer_to_probs(q: Question, ans: dict, raw: bool = False) -> np.ndarray: | |
| """SPEC answer -> probability vector in the record's option order (noul: [P(false), P(true)]).""" | |
| pre = "raw_" if raw else "" | |
| if q.type == "noul": | |
| p = float(ans[pre + "noul"]) | |
| v = np.array([1.0 - p, p]) | |
| elif q.type == "choice": | |
| d = ans[pre + "probabilities"] | |
| v = np.array([float(d.get(o.key, 0.0)) for o in q.options]) | |
| else: | |
| v = np.asarray(ans[pre + "probabilities"], dtype=float) | |
| if len(v) != len(q.options): | |
| raise ValueError(f"score answer has {len(v)} probabilities for {len(q.options)} levels") | |
| v = np.clip(v, 0.0, None) | |
| s = v.sum() | |
| return v / s if s > 0 else np.full(len(v), 1.0 / len(v)) | |
| def escalate_prob_of(ans: dict) -> Optional[float]: | |
| """Continuous escalate score if the answer exposes one (P(escalate), or 1 - P(argmax correct)).""" | |
| num = lambda v: isinstance(v, (int, float)) and not isinstance(v, bool) # noqa: E731 | |
| for k in _ESCALATE_PROB_KEYS: | |
| if num(ans.get(k)): | |
| return float(ans[k]) | |
| for k in _P_CORRECT_KEYS: | |
| if num(ans.get(k)): | |
| return 1.0 - float(ans[k]) | |
| return None | |
| def group_by_state(records: list[Record]) -> list[tuple[str, list[int]]]: | |
| """[(state, [record indices])] in first-seen order.""" | |
| groups: dict[str, list[int]] = {} | |
| for i, r in enumerate(records): | |
| groups.setdefault(r.state, []).append(i) | |
| return list(groups.items()) | |
| def _chunks(groups: list[tuple[str, list[int]]], budget: int) -> list[list[int]]: | |
| out, cur, n = [], [], 0 | |
| for gi, (_, idx) in enumerate(groups): | |
| if cur and n + len(idx) > budget: | |
| out.append(cur) | |
| cur, n = [], 0 | |
| cur.append(gi) | |
| n += len(idx) | |
| if cur: | |
| out.append(cur) | |
| return out | |
| def find_raw_hook(model) -> Optional[dict]: | |
| """kwargs that make predict_batch return uncalibrated probabilities, if the model has such a hook.""" | |
| try: | |
| params = inspect.signature(model.predict_batch).parameters | |
| except (TypeError, ValueError, AttributeError): | |
| return None | |
| for name, value in _RAW_KWARGS: | |
| if name in params: | |
| return {name: value} | |
| return None | |
| def run_predictions(model, records: list[Record], groups, publics, batch_questions: int = 32, | |
| kwargs: Optional[dict] = None): | |
| """Returns (answers per record or None, wall seconds, n forward calls, errors).""" | |
| kwargs = kwargs or {} | |
| answers: list[Optional[dict]] = [None] * len(records) | |
| errors: list[str] = [] | |
| n_calls = 0 | |
| def items_for(gis): | |
| return [(groups[g][0], {f"q{j}": publics[i] for j, i in enumerate(groups[g][1])}) for g in gis] | |
| def store(gis, outs): | |
| for g, res in zip(gis, outs): | |
| for j, i in enumerate(groups[g][1]): | |
| answers[i] = res["answers"][f"q{j}"] | |
| t0 = time.perf_counter() | |
| for gis in _chunks(groups, batch_questions): | |
| try: | |
| n_calls += 1 | |
| outs = model.predict_batch(items_for(gis), **kwargs) | |
| store(gis, outs) | |
| except Exception as e: # isolate the failing state(s) | |
| if len(gis) == 1: | |
| errors.append(f"{records[groups[gis[0]][1][0]].id}: {type(e).__name__}: {e}") | |
| continue | |
| for g in gis: | |
| try: | |
| n_calls += 1 | |
| store([g], model.predict_batch(items_for([g]), **kwargs)) | |
| except Exception as e2: | |
| errors.append(f"{records[groups[g][1][0]].id}: {type(e2).__name__}: {e2}") | |
| return answers, time.perf_counter() - t0, n_calls, errors | |
| # ---------------------------------------------------------------- metrics | |
| def reliability_bins(conf: np.ndarray, correct: np.ndarray, n_bins: int = N_BINS) -> list[dict]: | |
| b = np.minimum((conf * n_bins).astype(int), n_bins - 1) | |
| out = [] | |
| for k in range(n_bins): | |
| m = b == k | |
| out.append({"lo": k / n_bins, "hi": (k + 1) / n_bins, "count": int(m.sum()), | |
| "confidence": float(conf[m].mean()) if m.any() else None, | |
| "accuracy": float(correct[m].mean()) if m.any() else None}) | |
| return out | |
| def ece(conf: np.ndarray, correct: np.ndarray, n_bins: int = N_BINS) -> float: | |
| if len(conf) == 0: | |
| return float("nan") | |
| tot = 0.0 | |
| for b in reliability_bins(conf, correct, n_bins): | |
| if b["count"]: | |
| tot += b["count"] * abs(b["accuracy"] - b["confidence"]) | |
| return tot / len(conf) | |
| def prob_metrics(P: list[np.ndarray], y: np.ndarray, types: Optional[list[str]] = None, | |
| n_bins: int = N_BINS) -> dict: | |
| """Accuracy / NLL / Brier / ECE for per-question probability vectors P and gold indices y.""" | |
| n = len(P) | |
| if n == 0: | |
| return {"n": 0} | |
| pred = np.array([int(np.argmax(p)) for p in P]) | |
| conf = np.array([float(p.max()) for p in P]) | |
| p_gold = np.array([float(p[t]) for p, t in zip(P, y)]) | |
| correct = pred == y | |
| brier = np.array([float(((p - np.eye(len(p))[t]) ** 2).sum()) for p, t in zip(P, y)]) | |
| out = {"n": n, "accuracy": float(correct.mean()), "nll": float(-np.log(np.clip(p_gold, EPS, 1)).mean()), | |
| "brier": float(brier.mean()), "ece": float(ece(conf, correct, n_bins)), | |
| "mean_confidence": float(conf.mean())} | |
| if types is not None: | |
| sc = [i for i, t in enumerate(types) if t == "score"] | |
| if sc: | |
| out["score_mae_levels"] = float(np.mean([abs(pred[i] - y[i]) for i in sc])) | |
| out["score_rps"] = float(np.mean([_rps(P[i], y[i]) for i in sc])) | |
| return out | |
| def _rps(p: np.ndarray, t: int) -> float: | |
| k = len(p) | |
| cp = np.cumsum(p)[:-1] | |
| cy = (np.arange(k - 1) >= t).astype(float) | |
| return float(((cp - cy) ** 2).sum() / (k - 1)) | |
| def auroc(labels: np.ndarray, scores: np.ndarray) -> Optional[float]: | |
| labels = np.asarray(labels).astype(int) | |
| if len(labels) < 2 or labels.min() == labels.max(): | |
| return None | |
| from sklearn.metrics import roc_auc_score | |
| return float(roc_auc_score(labels, scores)) | |
| def selective_accuracy(correct: np.ndarray, order: np.ndarray, coverages=COVERAGES) -> dict: | |
| """`order`: record indices from most to least trusted. Accuracy on the top c fraction.""" | |
| n = len(order) | |
| out = {} | |
| for c in coverages: | |
| k = max(1, int(round(c * n))) | |
| out[f"{int(round(c * 100))}%"] = float(correct[order[:k]].mean()) if n else None | |
| return out | |
| def escalate_metrics(correct: np.ndarray, conf: np.ndarray, flags: Optional[np.ndarray], | |
| esc_prob: Optional[np.ndarray]) -> dict: | |
| err = ~correct | |
| out: dict[str, Any] = {"n": int(len(correct)), "error_rate": float(err.mean()) if len(err) else None} | |
| # baseline: max probability (MSP) | |
| out["confidence_baseline"] = {"auroc_error": auroc(err, -conf), | |
| "accuracy_at_coverage": selective_accuracy(correct, np.argsort(-conf, kind="stable"))} | |
| if esc_prob is not None: | |
| out["signal"] = "escalate probability" | |
| score = esc_prob | |
| elif flags is not None: | |
| out["signal"] = "escalate flag only (binary; ties broken by confidence)" | |
| score = flags.astype(float) | |
| else: | |
| out["signal"] = None | |
| return out | |
| order = np.lexsort((-conf, score)) # low escalate score first, then high confidence | |
| out["auroc_error"] = auroc(err, score) | |
| out["accuracy_at_coverage"] = selective_accuracy(correct, order) | |
| if flags is not None: | |
| f = flags.astype(bool) | |
| out["flag"] = {"escalate_rate": float(f.mean()), | |
| "accuracy_not_escalated": float(correct[~f].mean()) if (~f).any() else None, | |
| "accuracy_escalated": float(correct[f].mean()) if f.any() else None, | |
| "error_recall": float((f & err).sum() / err.sum()) if err.any() else None, | |
| "precision": float((f & err).sum() / f.sum()) if f.any() else None} | |
| return out | |
| def prior_baseline(records: list[Record], fit: list[Record]) -> list[np.ndarray]: | |
| counts: dict[tuple[str, str], Counter] = defaultdict(Counter) | |
| for r in fit: | |
| counts[(r.source, r.question.type)][r.label] += 1 | |
| out = [] | |
| for r in records: | |
| c = counts.get((r.source, r.question.type), Counter()) | |
| v = np.array([c.get(k, 0) + 1.0 for k in range(r.question.n_options)]) | |
| out.append(v / v.sum()) | |
| return out | |
| def uniform_baseline(records: list[Record]) -> dict: | |
| ns = np.array([r.question.n_options for r in records], dtype=float) | |
| if not len(ns): | |
| return {"n": 0} | |
| return {"n": len(ns), "accuracy": float((1 / ns).mean()), "nll": float(np.log(ns).mean()), | |
| "brier": float((1 - 1 / ns).mean()), "ece": 0.0, "mean_confidence": float((1 / ns).mean())} | |
| # ---------------------------------------------------------------- latency | |
| def measure_latency(model, items: list[tuple[str, dict]], n_calls: int = 30, warmup: int = 2) -> dict: | |
| if not items or n_calls <= 0: | |
| return {} | |
| for s, q in items[:warmup]: | |
| model.predict(s, q) | |
| ms, nq = [], [] | |
| for s, q in items[:n_calls]: | |
| t0 = time.perf_counter() | |
| model.predict(s, q) | |
| ms.append((time.perf_counter() - t0) * 1000) | |
| nq.append(len(q)) | |
| a = np.array(ms) | |
| return {"calls": len(ms), "p50_ms": float(np.percentile(a, 50)), "p95_ms": float(np.percentile(a, 95)), | |
| "mean_ms": float(a.mean()), "mean_questions_per_call": float(np.mean(nq))} | |
| # ---------------------------------------------------------------- plot | |
| _BLUE, _ORANGE, _INK, _INK2, _GRID = "#2a78d6", "#eb6834", "#0b0b0b", "#52514e", "#e4e3df" | |
| def plot_reliability(path: Path, panels: dict[str, dict[str, Optional[list[dict]]]], title: str) -> None: | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| names = [k for k, v in panels.items() if v.get("calibrated")] | |
| fig, axes = plt.subplots(2, len(names), figsize=(3.6 * len(names), 5.2), squeeze=False, | |
| gridspec_kw={"height_ratios": [3, 1], "hspace": 0.08, "wspace": 0.28}) | |
| series = [("raw", "raw (T=1)", _ORANGE), ("calibrated", "calibrated", _BLUE)] | |
| for col, name in enumerate(names): | |
| ax, axh = axes[0][col], axes[1][col] | |
| ax.plot([0, 1], [0, 1], ls="--", lw=1, color=_INK2, alpha=0.6, label="perfect") | |
| width = 1 / N_BINS | |
| for i, (key, label, color) in enumerate(series): | |
| bins = panels[name].get(key) | |
| if not bins: | |
| continue | |
| pts = [(b["confidence"], b["accuracy"]) for b in bins if b["count"]] | |
| ax.plot([p[0] for p in pts], [p[1] for p in pts], "-o", lw=2, ms=5, color=color, label=label, | |
| markeredgecolor="white", markeredgewidth=1) | |
| lo = np.array([b["lo"] for b in bins]) | |
| off = (i - 0.5) * width * 0.45 if panels[name].get("raw") else 0 | |
| axh.bar(lo + width / 2 + off, [b["count"] for b in bins], width=width * 0.42, color=color, | |
| edgecolor="white", linewidth=0.5) | |
| n = sum(b["count"] for b in panels[name]["calibrated"]) | |
| ax.set_title(f"{name} (n={n})", fontsize=10, color=_INK, loc="left") | |
| for a in (ax, axh): | |
| a.set_xlim(0, 1) | |
| a.grid(color=_GRID, lw=0.6) | |
| a.set_axisbelow(True) | |
| for s in ("top", "right"): | |
| a.spines[s].set_visible(False) | |
| for s in ("left", "bottom"): | |
| a.spines[s].set_color(_INK2) | |
| a.tick_params(colors=_INK2, labelsize=8) | |
| ax.set_ylim(0, 1) | |
| ax.set_xticklabels([]) | |
| axh.set_xlabel("confidence (max prob)", fontsize=9, color=_INK2) | |
| if col == 0: | |
| ax.set_ylabel("accuracy", fontsize=9, color=_INK2) | |
| axh.set_ylabel("questions", fontsize=9, color=_INK2) | |
| ax.legend(fontsize=8, frameon=False, loc="upper left") | |
| fig.suptitle(title, fontsize=11, color=_INK, x=0.01, ha="left") | |
| fig.savefig(path, dpi=120, bbox_inches="tight", facecolor="white") | |
| plt.close(fig) | |
| # ---------------------------------------------------------------- driver | |
| def evaluate(model, records: list[Record], *, name: str, out_dir: str | Path, | |
| record_files: Optional[list[str]] = None, checkpoint: Optional[str] = None, | |
| batch_questions: int = 32, latency_calls: int = 30, prior_records: Optional[list[Record]] = None, | |
| prior_source: str = "in-sample", raw: bool = True, seed: int = 0) -> dict: | |
| out_dir = Path(out_dir) | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| record_files = record_files or ["-"] * len(records) | |
| groups = group_by_state(records) | |
| publics = [question_to_public(r.question) for r in records] | |
| notes: list[str] = [] | |
| answers, wall_s, n_calls, errors = run_predictions(model, records, groups, publics, batch_questions) | |
| raw_answers, raw_mode, raw_in_answers = None, None, False | |
| if raw: | |
| hook = find_raw_hook(model) | |
| if hook is not None: | |
| raw_answers, _, _, raw_err = run_predictions(model, records, groups, publics, batch_questions, hook) | |
| raw_mode = f"predict_batch(..., {', '.join(f'{k}={v!r}' for k, v in hook.items())})" | |
| errors += [f"[raw] {e}" for e in raw_err] | |
| elif any(a is not None and ("raw_probabilities" in a or "raw_noul" in a) for a in answers): | |
| raw_answers, raw_mode, raw_in_answers = answers, "raw_* fields in the answers", True | |
| else: | |
| notes.append("The model exposes no uncalibrated output (no raw kwarg on predict_batch, no raw_* " | |
| "fields), so ECE is reported for calibrated probabilities only.") | |
| ok = [i for i, a in enumerate(answers) if a is not None] | |
| if len(ok) < len(records): | |
| notes.append(f"{len(records) - len(ok)} records failed to predict and are excluded (see `errors`).") | |
| recs = [records[i] for i in ok] | |
| y = np.array([r.label for r in recs]) | |
| types = [r.question.type for r in recs] | |
| sources = [r.source for r in recs] | |
| files = [record_files[i] for i in ok] | |
| P = [answer_to_probs(r.question, answers[i]) for r, i in zip(recs, ok)] | |
| R = None | |
| if raw_answers is not None: | |
| try: | |
| R = [answer_to_probs(r.question, raw_answers[i], raw=raw_in_answers) | |
| for r, i in zip(recs, ok)] | |
| except (KeyError, TypeError) as e: | |
| notes.append(f"raw probabilities unusable ({type(e).__name__}: {e}); calibrated only.") | |
| R, raw_mode = None, None | |
| prior_P = prior_baseline(recs, prior_records if prior_records is not None else recs) | |
| def block(idx: list[int]) -> dict: | |
| sub_t = [types[i] for i in idx] | |
| m = {"model": prob_metrics([P[i] for i in idx], y[idx], sub_t)} | |
| if R is not None: | |
| rm = prob_metrics([R[i] for i in idx], y[idx], sub_t) | |
| m["model"].update({"ece_raw": rm["ece"], "nll_raw": rm["nll"], "brier_raw": rm["brier"], | |
| "accuracy_raw": rm["accuracy"]}) | |
| m["prior"] = prob_metrics([prior_P[i] for i in idx], y[idx]) | |
| m["uniform"] = uniform_baseline([recs[i] for i in idx]) | |
| return m | |
| def by(keys: list[str]) -> dict: | |
| d: dict[str, list[int]] = defaultdict(list) | |
| for i, k in enumerate(keys): | |
| d[k].append(i) | |
| return {k: block(v) for k, v in sorted(d.items())} | |
| all_idx = list(range(len(recs))) | |
| report: dict[str, Any] = { | |
| "name": name, "checkpoint": checkpoint, "created": _dt.datetime.now().isoformat(timespec="seconds"), | |
| "model_name": _model_name(model), | |
| "max_len": getattr(model, "max_len", None), "long_state": getattr(model, "long_state", None), | |
| "device": str(getattr(model, "device", "unknown")), | |
| "files": sorted(set(record_files)), "n_records": len(records), "n_evaluated": len(recs), | |
| "n_states": len(groups), "raw_probabilities": raw_mode, "prior_fit": prior_source, | |
| "overall": block(all_idx), "by_type": by(types), "by_source": by(sources), "by_file": by(files), | |
| } | |
| # reliability | |
| conf = np.array([p.max() for p in P]) | |
| correct = np.array([int(np.argmax(p)) for p in P]) == y | |
| rconf = np.array([p.max() for p in R]) if R is not None else None | |
| rcorrect = (np.array([int(np.argmax(p)) for p in R]) == y) if R is not None else None | |
| panels: dict[str, dict] = {} | |
| for nm, idx in [("overall", all_idx)] + [(t, [i for i in all_idx if types[i] == t]) for t in ("choice", "score", "noul")]: | |
| if not idx: | |
| continue | |
| panels[nm] = {"calibrated": reliability_bins(conf[idx], correct[idx]), | |
| "raw": reliability_bins(rconf[idx], rcorrect[idx]) if R is not None else None} | |
| report["reliability"] = panels | |
| png = out_dir / f"{name}_reliability.png" | |
| if recs: | |
| plot_reliability(png, panels, f"Reliability - {name}") | |
| report["reliability_png"] = png.name | |
| # escalate head | |
| ans_ok = [answers[i] for i in ok] | |
| flags = np.array([bool(a.get("escalate")) for a in ans_ok]) if ans_ok and all("escalate" in a for a in ans_ok) else None | |
| eps = [escalate_prob_of(a) for a in ans_ok] | |
| esc_prob = np.array(eps, dtype=float) if ans_ok and all(e is not None for e in eps) else None | |
| report["escalate"] = escalate_metrics(correct, conf, flags, esc_prob) | |
| if esc_prob is None and flags is not None: | |
| notes.append("Answers carry only the boolean `escalate` flag, so escalate AUROC is that of a binary " | |
| "score; expose the escalate probability for a full ROC.") | |
| # speed | |
| rng = random.Random(seed) | |
| sample = rng.sample(range(len(groups)), min(len(groups), latency_calls + 2)) if latency_calls > 0 else [] | |
| items = [(groups[g][0], {f"q{j}": publics[i] for j, i in enumerate(groups[g][1])}) for g in sample] | |
| report["speed"] = {"single_call": measure_latency(model, items, latency_calls), | |
| "batched": {"questions": len(records), "states": len(groups), "forward_calls": n_calls, | |
| "batch_questions": batch_questions, "wall_s": wall_s, | |
| "questions_per_s": len(records) / wall_s if wall_s > 0 else None, | |
| "states_per_s": len(groups) / wall_s if wall_s > 0 else None}} | |
| report["errors"] = errors[:50] | |
| report["n_errors"] = len(errors) | |
| report["notes"] = notes | |
| (out_dir / f"{name}.json").write_text(json.dumps(_clean(report), indent=2, ensure_ascii=False)) | |
| (out_dir / f"{name}.md").write_text(render_markdown(report)) | |
| return report | |
| def _model_name(model) -> str: | |
| for a in ("name", "model_name"): | |
| v = getattr(model, a, None) | |
| if isinstance(v, str) and v: | |
| return v | |
| return type(model).__name__ | |
| def _clean(x): | |
| if isinstance(x, dict): | |
| return {k: _clean(v) for k, v in x.items()} | |
| if isinstance(x, (list, tuple)): | |
| return [_clean(v) for v in x] | |
| if isinstance(x, (np.floating, float)): | |
| return None if math.isnan(float(x)) else float(x) | |
| if isinstance(x, np.integer): | |
| return int(x) | |
| if isinstance(x, np.bool_): | |
| return bool(x) | |
| return x | |
| # ---------------------------------------------------------------- markdown | |
| def _f(x, pct=False, nd=4): | |
| if x is None or (isinstance(x, float) and math.isnan(x)): | |
| return "-" | |
| return f"{100 * x:.1f}%" if pct else f"{x:.{nd}f}" | |
| def _metric_rows(title: str, groups: dict[str, dict], raw: bool) -> list[str]: | |
| head = "| {} | n | acc | prior acc | uniform acc | NLL | Brier | ECE cal |{}".format( | |
| title, " ECE raw |" if raw else "") | |
| sep = "|" + "---|" * (8 + raw) | |
| rows = [head, sep] | |
| for k, g in groups.items(): | |
| m = g["model"] | |
| rows.append(f"| {k} | {m['n']} | {_f(m.get('accuracy'), True)} | {_f(g['prior'].get('accuracy'), True)} | " | |
| f"{_f(g['uniform'].get('accuracy'), True)} | {_f(m.get('nll'))} | {_f(m.get('brier'))} | " | |
| f"{_f(m.get('ece'))} |" + (f" {_f(m.get('ece_raw'))} |" if raw else "")) | |
| return rows | |
| def render_markdown(r: dict) -> str: | |
| raw = r.get("raw_probabilities") is not None | |
| o = r["overall"] | |
| L = [f"# Evaluation: {r['name']}", "", | |
| f"- checkpoint: `{r.get('checkpoint')}` model: `{r.get('model_name')}` device: `{r.get('device')}`", | |
| f"- data: {', '.join(f'`{f}`' for f in r['files'])}", | |
| f"- {r['n_evaluated']}/{r['n_records']} questions over {r['n_states']} distinct states; created {r['created']}", | |
| f"- raw (uncalibrated) probabilities: {r.get('raw_probabilities') or 'not available'}", | |
| *([f"- context: max_len {r['max_len']}, long states: {r.get('long_state') or 'truncate'}"] | |
| if isinstance(r.get("max_len"), int) else []), | |
| f"- prior baseline fitted on: {r['prior_fit']}", "", | |
| "## Overall", "", | |
| "| system | acc | NLL | Brier | ECE |", "|---|---|---|---|---|"] | |
| m = o["model"] | |
| L.append(f"| model (calibrated) | {_f(m.get('accuracy'), True)} | {_f(m.get('nll'))} | {_f(m.get('brier'))} | {_f(m.get('ece'))} |") | |
| if raw: | |
| L.append(f"| model (raw) | {_f(m.get('accuracy_raw'), True)} | {_f(m.get('nll_raw'))} | {_f(m.get('brier_raw'))} | {_f(m.get('ece_raw'))} |") | |
| for b, label in (("prior", "class prior / majority"), ("uniform", "uniform random (expected)")): | |
| mm = o[b] | |
| L.append(f"| {label} | {_f(mm.get('accuracy'), True)} | {_f(mm.get('nll'))} | {_f(mm.get('brier'))} | {_f(mm.get('ece'))} |") | |
| L += ["", "## By question type", ""] + _metric_rows("type", r["by_type"], raw) | |
| sc = r["by_type"].get("score", {}).get("model", {}) | |
| if "score_mae_levels" in sc: | |
| L += ["", f"Score questions: mean |argmax level - gold| = {sc['score_mae_levels']:.3f} levels, " | |
| f"RPS = {sc['score_rps']:.4f}."] | |
| L += ["", "## By source", ""] + _metric_rows("source", r["by_source"], raw) | |
| if len(r["by_file"]) > 1: | |
| L += ["", "## By file", ""] + _metric_rows("file", r["by_file"], raw) | |
| if r.get("reliability_png"): | |
| L += ["", "## Reliability", "", f"", | |
| "", f"{N_BINS} equal-width bins over top-label confidence; bars = questions per bin."] | |
| e = r["escalate"] | |
| L += ["", "## Escalate head", "", f"Signal: {e.get('signal') or 'none (answers carry no escalate field)'}. " | |
| f"Error rate: {_f(e.get('error_rate'), True)}.", "", | |
| "| ranking | AUROC (detect wrong answers) | acc @50% | acc @80% | acc @100% |", "|---|---|---|---|---|"] | |
| for label, d in (("escalate head", e), ("max-prob baseline", e["confidence_baseline"])): | |
| if d.get("accuracy_at_coverage"): | |
| c = d["accuracy_at_coverage"] | |
| L.append(f"| {label} | {_f(d.get('auroc_error'))} | {_f(c.get('50%'), True)} | {_f(c.get('80%'), True)} | {_f(c.get('100%'), True)} |") | |
| if e.get("flag"): | |
| fl = e["flag"] | |
| L += ["", f"`escalate=True` on {_f(fl['escalate_rate'], True)} of questions; accuracy when not escalated " | |
| f"{_f(fl['accuracy_not_escalated'], True)}, when escalated {_f(fl['accuracy_escalated'], True)}; " | |
| f"catches {_f(fl['error_recall'], True)} of errors (precision {_f(fl['precision'], True)})."] | |
| s = r["speed"] | |
| sc1, bt = s.get("single_call") or {}, s["batched"] | |
| L += ["", "## Speed", ""] | |
| if sc1: | |
| L.append(f"- single `predict` calls ({sc1['calls']} calls, {sc1['mean_questions_per_call']:.1f} questions/call): " | |
| f"p50 {sc1['p50_ms']:.1f} ms, p95 {sc1['p95_ms']:.1f} ms, mean {sc1['mean_ms']:.1f} ms") | |
| if bt.get("questions_per_s"): | |
| L.append(f"- batched `predict_batch` (<= {bt['batch_questions']} questions/call, {bt['forward_calls']} calls): " | |
| f"{bt['questions_per_s']:.1f} questions/s, {bt['states_per_s']:.1f} states/s ({bt['wall_s']:.2f} s total)") | |
| L += ["", "## Notes", "", | |
| "- Brier = sum over options of (p - y)^2 (noul uses 2 options). ECE: 15 equal-width bins, top-label.", | |
| "- prior = per-(source, type) label-index frequencies (majority class / position); uniform = expected values."] | |
| L += [f"- {n}" for n in r.get("notes", [])] | |
| if r.get("n_errors"): | |
| L += ["", f"### Errors ({r['n_errors']})", ""] + [f"- `{x}`" for x in r["errors"][:20]] | |
| return "\n".join(L) + "\n" | |
| # ---------------------------------------------------------------- CLI | |
| def load_model(checkpoint: str, device: Optional[str] = None, **opts): | |
| """SystemOne.load with optional overrides (max_len, long_state, chunk_agg; None = checkpoint default).""" | |
| from jevlike.predict import SystemOne | |
| kw = {k: v for k, v in opts.items() if v is not None} | |
| if device and "device" in inspect.signature(SystemOne.load).parameters: | |
| kw["device"] = device | |
| return SystemOne.load(checkpoint, **kw) | |
| def context_opts(args) -> dict: | |
| """Long-context CLI options that were actually given (so older loaders keep working).""" | |
| return {k: v for k, v in (("max_len", args.max_len), ("long_state", args.long_state), | |
| ("chunk_agg", args.chunk_agg)) if v is not None} | |
| def add_context_args(ap: argparse.ArgumentParser) -> None: | |
| ap.add_argument("--max-len", type=int, default=None, | |
| help="override the checkpoint's max_len (tokens, <= 8192; default: checkpoint value)") | |
| ap.add_argument("--long-state", choices=["truncate", "chunk"], default=None, | |
| help="states longer than max_len: cut in the middle (default) or score overlapping " | |
| "windows and pool them (see jevlike/predict.py)") | |
| ap.add_argument("--chunk-agg", default=None, | |
| help="pooling rule for --long-state chunk: auto (default: choice/score mean of " | |
| "log-probs, noul max), mean, max, linear, noisy_or") | |
| def main(argv=None) -> dict: | |
| ap = argparse.ArgumentParser(description="Evaluate a jevlike SystemOne checkpoint on Record JSONL files.") | |
| ap.add_argument("--checkpoint", required=True) | |
| ap.add_argument("--data", nargs="+", required=True, help="Record JSONL file(s)") | |
| ap.add_argument("--out", default="reports/") | |
| ap.add_argument("--name", default=None, help="report name (default: <checkpoint>__<data stems>)") | |
| ap.add_argument("--device", default=None, help="cpu / cuda (default: SystemOne auto)") | |
| ap.add_argument("--batch-questions", type=int, default=32, help="max questions per predict_batch call") | |
| ap.add_argument("--latency-calls", type=int, default=30, help="single predict calls to time (0 = skip)") | |
| ap.add_argument("--limit", type=int, default=None, help="use at most N records per file (first N)") | |
| ap.add_argument("--prior-data", default=None, help="fit the class-prior baseline on this JSONL (e.g. data/train.jsonl)") | |
| ap.add_argument("--no-raw", action="store_true", help="skip the uncalibrated pass") | |
| ap.add_argument("--seed", type=int, default=0) | |
| add_context_args(ap) | |
| args = ap.parse_args(argv) | |
| records, files = [], [] | |
| for f in args.data: | |
| rs = read_jsonl(f)[: args.limit] if args.limit else read_jsonl(f) | |
| records += rs | |
| files += [Path(f).name] * len(rs) | |
| prior = read_jsonl(args.prior_data) if args.prior_data else None | |
| name = args.name or f"{Path(args.checkpoint.rstrip('/')).name}__{'+'.join(Path(f).stem for f in args.data)}" | |
| model = load_model(args.checkpoint, args.device, **context_opts(args)) | |
| rep = evaluate(model, records, name=name, out_dir=args.out, record_files=files, checkpoint=args.checkpoint, | |
| batch_questions=args.batch_questions, latency_calls=args.latency_calls, prior_records=prior, | |
| prior_source=args.prior_data or "in-sample (evaluated records)", raw=not args.no_raw, | |
| seed=args.seed) | |
| m = rep["overall"]["model"] | |
| print(f"{name}: n={m['n']} acc={m['accuracy']:.4f} nll={m['nll']:.4f} brier={m['brier']:.4f} " | |
| f"ece={m['ece']:.4f}" + (f" ece_raw={m['ece_raw']:.4f}" if "ece_raw" in m else "")) | |
| print(f"wrote {Path(args.out) / (name + '.md')} and .json") | |
| return rep | |
| if __name__ == "__main__": | |
| main() | |