| """ |
| eval_help_shatter.py — the REAL eMoE metric (current-state §10.5 step 3). |
| |
| VAL 2.0889 says the hypernetwork GENERALIZES (held-out loss is low). It does NOT |
| say whether a minted adapter HELPS or SHATTERS on a novel request. This script |
| answers that, per task, by comparing — on the EXACT VAL loss construction |
| (train_hyper_sft.make_batch, output-only masking, ignore_index=-1) — the loss of |
| the gold output under: |
| |
| (a) the RIGHT-z minted adapter (hyper(z_task) -> set_deltas) |
| (b) the bare frozen base (set_deltas(None)) [z-independent] |
| (c) OPTIONAL a WRONG-z adapter (a neighbor task's z) [conditioning check] |
| |
| delta_help = loss_base - loss_adapt ( > 0 => adapter HELPS ) |
| ( < 0 => adapter SHATTERS ) |
| |
| It then buckets every (task, variant) point by its nearest-TRAIN-cluster cosine |
| distance — the SAME geometry the controller uses — and asks the question the |
| migration plan poses: *does the banked distance threshold separate help from |
| shatter?* Reported honestly. Per design (current-state §3) distance is NOVELTY, |
| not TRUST — shatter is meant to be CAUGHT by the verifier, not PREDICTED by |
| distance — so weak separation here is the expected, on-philosophy outcome, and |
| strong separation would be a (welcome) bonus, not a load-bearing assumption. |
| |
| WHY (task, variant) GRANULARITY: at serve time a single descriptor is encoded to |
| ONE z (the noisier "C" regime the taus were banked from). Each cached variant is |
| one such z, so per-(task,variant) is the serve-faithful unit. Per-task headline |
| numbers average over a task's variants. |
| |
| RUN (on the pod, with the v12 artifacts): |
| python eval_help_shatter.py \ |
| --base_ckpt ckpt_v12_190m_best.pt \ |
| --hyper_ckpt hyper_ckpt_v12.best.pt \ |
| --z_cache runs/hyper_v1/z_cache.pt \ |
| --tasks data/hyper_v1.jsonl \ |
| --val_frac 0.15 --seed 0 \ |
| --tau_rag_on 0.287 --tau_k_escalate 0.308 \ |
| --val_anchor 2.0889 \ |
| --wrong_z --out runs/hyper_v1/help_shatter.jsonl |
| |
| Logic-only self-test (no torch / no artifacts): python eval_help_shatter.py --self_test |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import math |
| import os |
| import sys |
| from statistics import mean, median |
|
|
|
|
| |
| |
| |
|
|
| def _auc(scores: list[float], labels: list[int]) -> float: |
| """AUC of `score` predicting label==1, via the rank-sum identity. 0.5 = no |
| signal. Used to ask: does nearest-cluster DISTANCE predict SHATTER?""" |
| pos = [s for s, y in zip(scores, labels) if y == 1] |
| neg = [s for s, y in zip(scores, labels) if y == 0] |
| if not pos or not neg: |
| return float("nan") |
| order = sorted(range(len(scores)), key=lambda i: scores[i]) |
| ranks = [0.0] * len(scores) |
| i = 0 |
| while i < len(order): |
| j = i |
| while j < len(order) and scores[order[j]] == scores[order[i]]: |
| j += 1 |
| avg = (i + j - 1) / 2.0 + 1.0 |
| for k in range(i, j): |
| ranks[order[k]] = avg |
| i = j |
| rank_sum_pos = sum(ranks[i] for i in range(len(scores)) if labels[i] == 1) |
| n_pos, n_neg = len(pos), len(neg) |
| return (rank_sum_pos - n_pos * (n_pos + 1) / 2.0) / (n_pos * n_neg) |
|
|
|
|
| def _bucket(d: float, tau_rag: float, tau_k: float) -> str: |
| if d < tau_rag: |
| return "in_dist" |
| if d < tau_k: |
| return "mid" |
| return "ood" |
|
|
|
|
| def summarize(records: list[dict], *, tau_rag: float, tau_k: float, |
| shatter_eps: float, val_anchor: float | None) -> dict: |
| """records: per-(task,variant) dicts with keys |
| task_id, variant, loss_base, loss_adapt, delta, dist.""" |
| n = len(records) |
| deltas = [r["delta"] for r in records] |
| helps = [1 if d > 0 else 0 for r, d in zip(records, deltas)] |
| shatters = [1 if d < -shatter_eps else 0 for d in deltas] |
|
|
| |
| by_task: dict[str, list[float]] = {} |
| base_by_task: dict[str, float] = {} |
| for r in records: |
| by_task.setdefault(r["task_id"], []).append(r["delta"]) |
| base_by_task[r["task_id"]] = r["loss_base"] |
| task_delta = {t: mean(v) for t, v in by_task.items()} |
| task_help = sum(1 for v in task_delta.values() if v > 0) |
| task_shatter = sum(1 for v in task_delta.values() if v < -shatter_eps) |
| n_tasks = len(task_delta) |
|
|
| out = { |
| "n_points": n, |
| "n_tasks": n_tasks, |
| "mean_loss_base": mean(r["loss_base"] for r in records), |
| "mean_loss_adapt": mean(r["loss_adapt"] for r in records), |
| "mean_delta": mean(deltas), |
| "median_delta": median(deltas), |
| "help_rate_points": sum(helps) / n if n else float("nan"), |
| "shatter_rate_points": sum(shatters) / n if n else float("nan"), |
| "help_rate_tasks": task_help / n_tasks if n_tasks else float("nan"), |
| "shatter_rate_tasks": task_shatter / n_tasks if n_tasks else float("nan"), |
| "shatter_eps": shatter_eps, |
| } |
| if val_anchor is not None: |
| out["val_anchor"] = val_anchor |
| out["adapt_vs_anchor"] = out["mean_loss_adapt"] - val_anchor |
|
|
| |
| buckets: dict[str, list[dict]] = {"in_dist": [], "mid": [], "ood": []} |
| for r in records: |
| buckets[_bucket(r["dist"], tau_rag, tau_k)].append(r) |
| out["by_bucket"] = {} |
| for name, rs in buckets.items(): |
| if not rs: |
| out["by_bucket"][name] = {"n": 0} |
| continue |
| ds = [r["delta"] for r in rs] |
| out["by_bucket"][name] = { |
| "n": len(rs), |
| "mean_delta": mean(ds), |
| "help_rate": sum(1 for d in ds if d > 0) / len(ds), |
| "shatter_rate": sum(1 for d in ds if d < -shatter_eps) / len(ds), |
| "mean_dist": mean(r["dist"] for r in rs), |
| } |
|
|
| |
| dist_all = [r["dist"] for r in records] |
| out["dist_predicts_shatter_auc"] = _auc(dist_all, shatters) |
| sh_d = [r["dist"] for r, s in zip(records, shatters) if s] |
| hp_d = [r["dist"] for r, s in zip(records, shatters) if not s] |
| out["mean_dist_shatter"] = mean(sh_d) if sh_d else float("nan") |
| out["mean_dist_nonshatter"] = mean(hp_d) if hp_d else float("nan") |
| return out |
|
|
|
|
| def print_report(s: dict, *, tau_rag: float, tau_k: float) -> None: |
| p = lambda *a: print(*a) |
| p("\n" + "=" * 70) |
| p("HELP vs SHATTER — right-z adapted loss vs bare base (VAL construction)") |
| p("=" * 70) |
| p(f"points (task,variant): {s['n_points']} tasks: {s['n_tasks']}") |
| p(f"mean loss base={s['mean_loss_base']:.4f} adapt={s['mean_loss_adapt']:.4f}" |
| f" (delta {s['mean_delta']:+.4f}, median {s['median_delta']:+.4f})") |
| if "val_anchor" in s: |
| flag = " <-- WARN: >0.3 off, config may differ from the trained run" \ |
| if abs(s["adapt_vs_anchor"]) > 0.30 else "" |
| p(f"anchor: VAL floor {s['val_anchor']:.4f}; mean adapt is " |
| f"{s['adapt_vs_anchor']:+.4f} vs floor{flag}") |
| p(f"\nHELP rate tasks={s['help_rate_tasks']*100:5.1f}% " |
| f"points={s['help_rate_points']*100:5.1f}% (delta > 0)") |
| p(f"SHATTER rate tasks={s['shatter_rate_tasks']*100:5.1f}% " |
| f"points={s['shatter_rate_points']*100:5.1f}% (delta < -{s['shatter_eps']})") |
| if "conditioning_win_rate" in s: |
| p(f"\nCONDITIONING (right-z beats wrong-z): " |
| f"{s['conditioning_win_rate']*100:5.1f}% of tasks " |
| f"(mean gap {s['conditioning_mean_gap']:+.4f} nats)") |
| p(f"\nby nearest-cluster distance bucket " |
| f"(in<{tau_rag} | mid<{tau_k} | ood>={tau_k}):") |
| p(f" {'bucket':<9}{'n':>6}{'mean_d':>9}{'help%':>8}{'shatter%':>10}{'mean_delta':>12}") |
| for name in ("in_dist", "mid", "ood"): |
| b = s["by_bucket"][name] |
| if b["n"] == 0: |
| p(f" {name:<9}{0:>6}{'—':>9}{'—':>8}{'—':>10}{'—':>12}") |
| continue |
| p(f" {name:<9}{b['n']:>6}{b['mean_dist']:>9.3f}" |
| f"{b['help_rate']*100:>7.1f}%{b['shatter_rate']*100:>9.1f}%" |
| f"{b['mean_delta']:>+12.4f}") |
| auc = s["dist_predicts_shatter_auc"] |
| p(f"\nDOES DISTANCE PREDICT SHATTER? AUC = {auc:.3f} " |
| f"(0.5 = no signal; mean dist shatter={s['mean_dist_shatter']:.3f} " |
| f"vs non-shatter={s['mean_dist_nonshatter']:.3f})") |
| if not math.isnan(auc): |
| if auc < 0.60: |
| p(" -> WEAK/NONE. Consistent with the design: distance is novelty, not") |
| p(" trust. Shatter must be CAUGHT by the verifier, not gated on distance.") |
| p(" The banked taus stay justified as RAG/k (novelty) knobs, NOT as a") |
| p(" shatter gate. Do NOT re-couple distance->alpha on this.") |
| else: |
| p(" -> Some separation. A BONUS signal, but per current-state §3 it enters") |
| p(" as a NEW knob only after replication; it does not silently gate alpha.") |
| p("=" * 70 + "\n") |
|
|
|
|
| |
| |
| |
|
|
| def run_eval(args) -> None: |
| import numpy as np |
| import torch |
| import tok_v9 |
| from train_hyper_sft import make_batch, load_tasks |
| from runtime_adapters import (HyperExpertRunner, GenConfig, build_cluster_index, |
| _replicate_split, _unit) |
|
|
| tok = tok_v9.build() |
| runner = HyperExpertRunner(args.base_ckpt, args.hyper_ckpt, tok, |
| device=args.device, gen=GenConfig()) |
| adapted, hyper = runner.adapted, runner.hyper |
| block, dev, d_z = runner.block_size, runner.device, runner.d_z |
|
|
| |
| tasks = load_tasks(args.tasks) |
| _, val_tasks = _replicate_split(tasks, args.val_frac, args.seed) |
| z_cache = torch.load(args.z_cache, map_location="cpu") |
|
|
| |
| clusters, cinfo = build_cluster_index(args.tasks, args.z_cache, |
| val_frac=args.val_frac, seed=args.seed, |
| scope="train") |
| print(f"[eval] {len(val_tasks)} val tasks | {cinfo['n_clusters']} train clusters " |
| f"| d_z {d_z} | device {dev}") |
|
|
| |
| def mean_z(tid): |
| Z = np.asarray(z_cache[tid].float().cpu().numpy(), dtype=np.float32) |
| if Z.ndim == 1: |
| Z = Z[None, :] |
| return _unit(Z.mean(0)) |
|
|
| records: list[dict] = [] |
| cond_gaps: list[float] = [] |
| skipped = 0 |
| val_ids = [t["task_id"] for t in val_tasks] |
|
|
| with torch.no_grad(): |
| for i, t in enumerate(val_tasks): |
| tid = t["task_id"] |
| examples = t.get("eval_examples") or t.get("train_examples") |
| if not examples or tid not in z_cache: |
| skipped += 1 |
| continue |
| b = make_batch(examples, tok, block, dev) |
| if b is None: |
| skipped += 1 |
| continue |
|
|
| |
| adapted.set_deltas(None) |
| loss_base = float(adapted(*b)[1].item()) |
|
|
| |
| Z = z_cache[tid] |
| Z = Z[None, :] if Z.ndim == 1 else Z |
| n_var = Z.shape[0] if args.max_variants <= 0 else min(args.max_variants, Z.shape[0]) |
| adapt_losses = [] |
| for v in range(n_var): |
| z = Z[v].float().to(dev) |
| dist_v = float(clusters.nearest(z.cpu().numpy())[0]) |
| adapted.set_deltas(hyper(z)) |
| lv = float(adapted(*b)[1].item()) |
| adapt_losses.append(lv) |
| records.append({"task_id": tid, "variant": v, |
| "loss_base": loss_base, "loss_adapt": lv, |
| "delta": loss_base - lv, "dist": dist_v}) |
|
|
| |
| if args.wrong_z and len(val_ids) > 1: |
| wrong_tid = val_ids[(i + 1) % len(val_ids)] |
| if wrong_tid in z_cache: |
| zw = torch.as_tensor(mean_z(wrong_tid), device=dev) |
| adapted.set_deltas(hyper(zw)) |
| loss_wrong = float(adapted(*b)[1].item()) |
| cond_gaps.append(loss_wrong - mean(adapt_losses)) |
|
|
| if (i + 1) % 25 == 0: |
| print(f" [{i+1}/{len(val_tasks)}] last base={loss_base:.3f} " |
| f"adapt={mean(adapt_losses):.3f}", flush=True) |
|
|
| if not records: |
| raise SystemExit("no records — check that z_cache keys match val task_ids " |
| "and tasks have eval_examples/train_examples.") |
| if skipped: |
| print(f"[eval] skipped {skipped} val tasks (no examples / not in cache)") |
|
|
| s = summarize(records, tau_rag=args.tau_rag_on, tau_k=args.tau_k_escalate, |
| shatter_eps=args.shatter_eps, val_anchor=args.val_anchor) |
| if cond_gaps: |
| s["conditioning_win_rate"] = sum(1 for g in cond_gaps if g > 0) / len(cond_gaps) |
| s["conditioning_mean_gap"] = mean(cond_gaps) |
|
|
| print_report(s, tau_rag=args.tau_rag_on, tau_k=args.tau_k_escalate) |
|
|
| if args.out: |
| os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True) |
| with open(args.out, "w") as f: |
| for r in records: |
| f.write(json.dumps(r) + "\n") |
| with open(args.out.rsplit(".", 1)[0] + ".summary.json", "w") as f: |
| json.dump(s, f, indent=2) |
| print(f"[eval] wrote {len(records)} records -> {args.out}") |
| print(f"[eval] wrote summary -> {args.out.rsplit('.', 1)[0]}.summary.json") |
|
|
|
|
| |
| |
| |
|
|
| def _self_test() -> None: |
| print("== eval_help_shatter self-test (pure python) ==") |
| res = [] |
|
|
| |
| assert abs(_auc([1, 2, 3, 4], [0, 0, 1, 1]) - 1.0) < 1e-9 |
| assert abs(_auc([1, 2, 3, 4], [1, 1, 0, 0]) - 0.0) < 1e-9 |
| assert abs(_auc([1, 1, 2, 2], [0, 1, 0, 1]) - 0.5) < 1e-9 |
| res.append(("auc", "1.0 / 0.0 / 0.5")) |
|
|
| |
| assert _bucket(0.20, 0.287, 0.308) == "in_dist" |
| assert _bucket(0.287, 0.287, 0.308) == "mid" |
| assert _bucket(0.308, 0.287, 0.308) == "ood" |
| res.append(("bucket boundaries", "half-open at tau")) |
|
|
| |
| |
| recs = [] |
| for k in range(40): |
| d = 0.10 + 0.005 * k |
| |
| delta = -0.3 if k % 9 == 0 else 0.4 |
| recs.append({"task_id": f"t{k}", "variant": 0, |
| "loss_base": 3.0, "loss_adapt": 3.0 - delta, |
| "delta": delta, "dist": d}) |
| s = summarize(recs, tau_rag=0.287, tau_k=0.308, shatter_eps=0.05, val_anchor=2.6) |
| assert s["n_points"] == 40 and s["n_tasks"] == 40 |
| assert s["help_rate_points"] > 0.8 and s["shatter_rate_points"] > 0.0 |
| assert 0.30 < s["dist_predicts_shatter_auc"] < 0.70, s["dist_predicts_shatter_auc"] |
| assert s["mean_loss_base"] == 3.0 |
| assert abs(s["adapt_vs_anchor"] - (s["mean_loss_adapt"] - 2.6)) < 1e-9 |
| res.append(("summarize uncorrelated", f"AUC={s['dist_predicts_shatter_auc']:.2f} " |
| f"help={s['help_rate_points']:.2f}")) |
|
|
| |
| recs2 = [] |
| for k in range(40): |
| d = 0.10 + 0.006 * k |
| delta = -0.3 if d > 0.30 else 0.4 |
| recs2.append({"task_id": f"t{k}", "variant": 0, |
| "loss_base": 3.0, "loss_adapt": 3.0 - delta, |
| "delta": delta, "dist": d}) |
| s2 = summarize(recs2, tau_rag=0.287, tau_k=0.308, shatter_eps=0.05, val_anchor=None) |
| assert s2["dist_predicts_shatter_auc"] > 0.85, s2["dist_predicts_shatter_auc"] |
| assert s2["by_bucket"]["ood"]["shatter_rate"] > s2["by_bucket"]["in_dist"]["shatter_rate"] |
| res.append(("summarize correlated", f"AUC={s2['dist_predicts_shatter_auc']:.2f}")) |
|
|
| print() |
| for k, v in res: |
| print(f" [ok ] {k:<26} -> {v}") |
| print(f"ALL {len(res)}/{len(res)} LOGIC TESTS PASSED.") |
| print("(the torch/artifact path runs on the pod via run_eval)") |
|
|
|
|
| def main() -> None: |
| ap = argparse.ArgumentParser(description="eMoE help-vs-shatter eval (step 3)") |
| ap.add_argument("--self_test", action="store_true") |
| ap.add_argument("--base_ckpt", default="ckpt_v12_190m_best.pt") |
| ap.add_argument("--hyper_ckpt", default="hyper_ckpt_v12.best.pt") |
| ap.add_argument("--z_cache", default="runs/hyper_v1/z_cache.pt") |
| ap.add_argument("--tasks", default="data/hyper_v1.jsonl") |
| ap.add_argument("--val_frac", type=float, default=0.15) |
| ap.add_argument("--seed", type=int, default=0) |
| ap.add_argument("--device", default=None) |
| ap.add_argument("--max_variants", type=int, default=0, |
| help="cap cached z variants per task (0 = all)") |
| ap.add_argument("--wrong_z", action="store_true", |
| help="also measure a neighbor-task's z (conditioning check)") |
| ap.add_argument("--tau_rag_on", type=float, default=0.287, |
| help="banked from serve_emoe.py --geometry (C) p90") |
| ap.add_argument("--tau_k_escalate", type=float, default=0.308, |
| help="banked from serve_emoe.py --geometry (C) p95") |
| ap.add_argument("--shatter_eps", type=float, default=0.05, |
| help="delta below -eps counts as shatter (nats)") |
| ap.add_argument("--val_anchor", type=float, default=None, |
| help="VAL floor to sanity-check mean adapt loss against (e.g. 2.0889)") |
| ap.add_argument("--out", default="") |
| args = ap.parse_args() |
|
|
| if args.self_test: |
| _self_test() |
| return |
| run_eval(args) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|