V12Emoe / eval_help_shatter.py
Daxamite's picture
Upload eval_help_shatter.py
bae7522 verified
Raw
History Blame Contribute Delete
19 kB
"""
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()