""" 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 # --------------------------------------------------------------------------- # Pure-python aggregation / separation logic (self-testable without torch) # --------------------------------------------------------------------------- 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): # average ranks within ties 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" # below tau_rag_on: controller trusts the mint, RAG off, k=1 if d < tau_k: return "mid" # RAG fires, still k=1 return "ood" # >= tau_k_escalate: RAG + escalation 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] # per-task headline (mean delta over a task's variants) 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 # by-distance buckets 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), } # THE migration question: does distance predict shatter? 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") # --------------------------------------------------------------------------- # Heavy path — runs on the pod with torch + the real artifacts # --------------------------------------------------------------------------- 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 # exact trainer split -> the val tasks the hyper NEVER trained on 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") # controller geometry: nearest TRAIN cluster (scope='train' == §10 calibration) 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}") # mean-z per val task (for the optional wrong-z neighbor pairing) 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) # (X, Y), output-only mask if b is None: skipped += 1 continue # (b) bare base — z-independent, compute once adapted.set_deltas(None) loss_base = float(adapted(*b)[1].item()) # (a) right-z, per cached variant (serve-faithful single-descriptor z) 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}) # (c) wrong-z conditioning check: a NEIGHBOR task's mean z 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)) # >0 => right better 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") # --------------------------------------------------------------------------- # Self-test — aggregation + separation logic only (no torch / artifacts) # --------------------------------------------------------------------------- def _self_test() -> None: print("== eval_help_shatter self-test (pure python) ==") res = [] # AUC sanity: perfectly separable, anti-separable, random 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 # ties -> chance res.append(("auc", "1.0 / 0.0 / 0.5")) # bucketing boundaries are half-open as documented assert _bucket(0.20, 0.287, 0.308) == "in_dist" assert _bucket(0.287, 0.287, 0.308) == "mid" # >= tau_rag_on assert _bucket(0.308, 0.287, 0.308) == "ood" # >= tau_k_escalate res.append(("bucket boundaries", "half-open at tau")) # summarize: a help-heavy set with a few shatters, distance UNCORRELATED # with shatter (the on-design case) -> AUC ~ 0.5, help rate high. recs = [] for k in range(40): d = 0.10 + 0.005 * k # spread of distances # most help (+0.4), every 9th shatters (-0.3) regardless of distance 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}")) # summarize: distance DOES predict shatter (far -> shatter) -> AUC high recs2 = [] for k in range(40): d = 0.10 + 0.006 * k delta = -0.3 if d > 0.30 else 0.4 # shatter only when far 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()