route2 instrumentation + lr fix + deploy overlay (route2-act divergence)

route2-act diverged (run 43): 33M kaiming A_q/B_q at delta_S's lr=3e-3 blew up
(gn 0.3->7.5 step 8, generations -> token salad, lp_t -11). Fixes:
- #167 separate quarantine lr (route2_quar_lr_scale=0.1) so the 60x-bigger fresh
  LoRA isn't trained at the main-knob lr.
- #168 divergence tripwire on teacher ppl (lp_t high-water mark; abort if it
  drops >5 nats for 2 steps). Relative so tiny-random smoke (flat lp_t~-11.9)
  doesn't false-trip.
- #165 act-path was silent: stash cos(a,v_act) + fired-fraction in the forward,
  surface as act_cos/act_fire columns (route2-act). smoke shows act_fire=0.64 =>
  the cos>0 sign test over-routes (fires on most tokens, not just hack ones).
- #166 print last train generation before FINAL EVAL (coherence eyeball).
- route2 v_act/v_grad refresh was firing but silent -- now announced.
- #162 plot_deploy_overlay.py: per-mode DEPLOY overlay from per_mode_deploy.json
  (honest shipped-model numbers, route2-safe). just plot-deploy.
- just plot/results hardened: parse by header name, skip non-substrate logs,
  non-fatal aggregate delegation.

Co-Authored-By: Claudypoo <288921227+claudypoo@users.noreply.github.com>
This commit is contained in:
wassname
2026-05-31 23:16:39 +00:00
co-authored by Claudypoo
parent ad048e59c6
commit 11bcdd2fe6
5 changed files with 247 additions and 18 deletions
+5
View File
@@ -133,6 +133,11 @@ def _delta_hook(layer: nn.Linear, args: tuple, y: Tensor) -> Tensor:
v_act = layer._antipasto_v_act.to(a.dtype) # [r] unit, hack-ward, in Vh coords (fp32 buffer -> a.dtype)
cos = (a @ v_act) / (a.norm(dim=-1).clamp_min(1e-6) * v_act.norm().clamp_min(1e-6))
m = cos > 0 # [...] bool
# Stash routing intensity so train.py can log it (else the act path is silent
# and over-routing -- m firing on ~half of all tokens, not just hack tokens --
# is invisible). fired = fraction of token positions routed to the quarantine.
layer._antipasto_act_fired = m.float().mean().detach()
layer._antipasto_act_cos = cos.mean().detach()
kept = torch.where(m.unsqueeze(-1), kept.detach(), kept)
return y + (kept + quar).to(y.dtype)
+88 -3
View File
@@ -55,6 +55,7 @@ from __future__ import annotations
import gzip
import json
import math
import os
import sys
import random
@@ -154,6 +155,11 @@ class Config:
# detach, single pass. "grad" (Arm A): per-rollout cos(g_b, v_grad) from a gate
# probe, routes by subtracting flagged rollouts from delta_S.grad post-backward.
route2_mask: Literal["act", "grad"] = "act"
# route2-only: the quarantine A_q/B_q (33M fresh kaiming params) is ~60x larger
# than delta_S (0.5M) and at the shared delta_S lr it diverged -- gn 0.3->7.5 at
# step 8, generations -> token salad, lp_t -11 (run 43). Give it its own lower lr.
# Scale of main lr; 1.0 = old (diverging) behaviour, 0.1 = the fix.
route2_quar_lr_scale: float = 0.1
# Scale-dependent knobs — every preset must set these to a real value;
# subclasses below override the defaults.
model: str = "Qwen/Qwen3-4B"
@@ -687,7 +693,17 @@ class StepLogger:
_Col("cos_post", 6, "cout", ".2f", "hack-ward fraction AFTER projection (want ~0: all removed)"),
_Col("fired", 5, "fired", ".2f", "fraction of modules where projection fired"),
]
if arm == "routing":
# route2 act-mask: no v_hack grad projection, but the forward routes by
# cos(activation, v_act)>0. Surface that routing intensity (reuses the row's
# cos_pre/fired keys, populated from the stashed act stats in train.py) so the
# act path is no longer silent -- watch `fired` for over-routing (>>0.5 means
# the sign test fires on generic tokens, starving delta_S onto the quarantine).
if arm == "routing2_act":
cols += [
_Col("cos_pre", 7, "act_cos", "+.2f", "mean cos(activation, v_act): forward routing alignment"),
_Col("fired", 6, "act_fire", ".2f", "fraction of token positions routed to quarantine (cos>0)"),
]
if arm in ("routing", "routing2_act", "routing2_grad"):
cols += [
_Col("hack_deploy", 7, "hk_dep", "+.2f", "DEPLOY-eval hack (quarantine deleted = deployed model)"),
_Col("solve_deploy", 7, "slv_dep", "+.2f", "DEPLOY-eval solve"),
@@ -754,6 +770,7 @@ def main(cfg: Config) -> int:
is_route2 = cfg.intervention == "route2"
is_route2_grad = is_route2 and cfg.route2_mask == "grad"
is_route2_act = is_route2 and cfg.route2_mask == "act"
wrappers = wrap_model_with_antipasto(
model, model_name, CACHE_ROOT, device,
quarantine_rank=cfg.route2_quarantine_rank if is_route2 else None,
@@ -924,10 +941,18 @@ def main(cfg: Config) -> int:
f"G_s={G_s} student + G_t={G_t} teacher per prompt (mix_ratio={cfg.mix_ratio})."
)
# Quarantine (A_q/B_q) gets its own lower lr: it is ~60x bigger than delta_S and
# freshly kaiming-init, so the shared lr diverged it (run 43). Separate param group
# so the scheduler scales both proportionally (the group's lr rides on `lr` via the
# ratio captured here -- LinearLR/CosineAnnealingLR multiply each group's base lr).
quar_lr = lr * cfg.route2_quar_lr_scale
opt = torch.optim.AdamW(
delta_params + delta_hack_params + quar_params, lr=lr, weight_decay=cfg.weight_decay,
betas=(adam_beta1, adam_beta2),
[{"params": delta_params + delta_hack_params, "lr": lr},
{"params": quar_params, "lr": quar_lr}],
lr=lr, weight_decay=cfg.weight_decay, betas=(adam_beta1, adam_beta2),
)
if quar_params:
logger.info(f"route2 quarantine lr = {quar_lr:.1e} ({cfg.route2_quar_lr_scale}x main lr {lr:.1e})")
# Linear warmup over `warmup_frac * steps`, then cosine decay to 0 over the rest.
# Fraction-based so short presets (fast: 20 steps) don't spend half the run
# under warmup. Canonical full-preset: 0.1 * 100 = 10 (matches ariahw config.py:141).
@@ -1055,6 +1080,17 @@ def main(cfg: Config) -> int:
rollout_log_path.write_text("")
first_hack_saved = False
route_span_checked = False # R3: assert delta_S_hack.grad in span(V) once
last_gen_sample = None # first student rollout of the latest step (for collapse inspection)
diverged_steps = 0 # consecutive steps with collapsed teacher ppl (divergence tripwire)
lp_t_best = -float("inf") # coherence high-water mark (best teacher gen_logp seen)
# ppl_t = exp(-lp_t) on the FIXED teacher rollouts is a free coherence gauge.
# Divergence is a DROP from the run's own best coherence, not an absolute level:
# a real model sits at lp_t ~ -0.7 and craters to -11..-21 when it diverges (run
# 43: lr too high on the 33M quarantine, generations -> token salad), a ~10-nat
# drop. A relative threshold also keeps `just smoke` green -- the tiny-random model
# has an intrinsic lp_t ~ -11.9 (uniform logp) but it stays flat, so it never DROPS.
# Abort if lp_t falls this far below its best for 2 steps running (advantage dead).
DIVERGENCE_DROP = 5.0 # nats below best (e^5 ~ 150x worse ppl); never in healthy runs
dumped_hack_classes: set[str] = set() # first full example of each hack class -> verbose log
teacher_dumped = False
# Per-mode learning tracker (the substrate UAT: did the student learn EACH hack,
@@ -1503,6 +1539,18 @@ def main(cfg: Config) -> int:
diag = {"mean_cos_pre": float("nan"), "mean_cos_post": float("nan"),
"frac_fired": float("nan"), "mean_cos_pre_s": float("nan"),
"mean_cos_pre_t": float("nan")}
# route2 act-mask: the forward stashed per-layer fired-fraction + mean cos
# (cos(a,v_act)). Surface them in cin (mean cos) and fired (routed fraction)
# so over-routing is visible -- a frozen sign-test direction fires on ~half
# of all tokens, starving delta_S and dumping learning onto the quarantine.
if is_route2_act:
fired = [info["layer"]._antipasto_act_fired for info in wrappers.values()
if hasattr(info["layer"], "_antipasto_act_fired")]
coss = [info["layer"]._antipasto_act_cos for info in wrappers.values()
if hasattr(info["layer"], "_antipasto_act_cos")]
if fired:
diag["frac_fired"] = float(torch.stack(fired).mean())
diag["mean_cos_pre"] = float(torch.stack(coss).mean())
# route2 grad-mask: report the mean per-module per-rollout flag rate so
# we can watch the mask actually fire (and rise as hacks emerge).
if is_route2_grad and step_flagged:
@@ -1595,6 +1643,13 @@ def main(cfg: Config) -> int:
if _was_training:
model.train()
refr = f"route2:{cfg.route2_mask}"
# Announce it -- the route2 refresh was previously silent (only the
# v_hack path logged "refresh@step"), so it looked like the mask never
# refreshed. NOTE: this fires AFTER opt.step(), so if the model is
# already diverging the re-extracted direction is extracted on a broken
# model -- watch lp_t / ppl_t around the refresh step.
logger.info(f"route2 {cfg.route2_mask}-mask refreshed@step{step} "
f"({len(wrappers)} modules, quarantine ablated during extract)")
if v_hack is not None and do_refresh:
from .extract_vhack_grad import extract_v_hack
if cfg.vhack_pairs_path is not None:
@@ -1810,6 +1865,25 @@ def main(cfg: Config) -> int:
with rollout_log_path.open("a") as fh:
for rec in step_rollouts:
fh.write(json.dumps(rec) + "\n")
if step_rollouts:
last_gen_sample = (step, step_rollouts[0]) # newest student gen for the final dump
# Divergence tripwire on teacher perplexity (free coherence gauge, see init).
ppl_t = math.exp(-lp_t_mean) if math.isfinite(lp_t_mean) else float("inf")
if math.isfinite(lp_t_mean):
lp_t_best = max(lp_t_best, lp_t_mean)
diverged = math.isfinite(lp_t_mean) and lp_t_mean < lp_t_best - DIVERGENCE_DROP
diverged_steps = diverged_steps + 1 if diverged else 0
if diverged_steps >= 2:
logger.error(
f"DIVERGED at step {step}: lp_t={lp_t_mean:.1f} (ppl_t={ppl_t:.0e}), {lp_t_best - lp_t_mean:.1f} "
f"nats below best {lp_t_best:.1f}, for {diverged_steps} steps -- policy collapsed "
f"(gn={gn:.1f}). Aborting to save GPU. Likely lr too high (route2: lower --route2-quar-lr-scale).")
if last_gen_sample:
_s, _r = last_gen_sample
logger.error(f"--- last student gen (step {_s}, reward={_r['reward']:+.2f}) ---\n"
f"{_r['text'][:800]}\n--- END (token salad => divergence confirmed) ---")
raise RuntimeError(f"training diverged (ppl_t={ppl_t:.0e} at step {step})")
if (step + 1) % 25 == 0:
save_ckpt(rows) # survive early kills; ~12 days for the full sweep
if not first_hack_saved and hack_s_n > 0:
@@ -1881,6 +1955,17 @@ def main(cfg: Config) -> int:
logger.info(f"||B_q|| sum = {bq_norm:.4f} (SHOULD: >0; ELSE quarantine never seeded)")
assert bq_norm > 0.0, "route2: B_q never moved -> quarantine never seeded (mask never fired?)"
# Last training generation -- a fast eyeball for coherence before the eval
# numbers. SHOULD: real code/prose for the problem. If it is token salad the
# policy diverged and every eval number below is meaningless (see ppl_t / lp_t).
if last_gen_sample is not None:
_s, _r = last_gen_sample
logger.info(
f"\n\n=== LAST TRAIN GEN (step {_s}, reward={_r['reward']:+.2f}, "
f"gt_pass={_r['gt_pass']}, hacked={_r['hacked_E']}) ===\n"
f"SHOULD: coherent code/prose. ELSE token salad => diverged, eval below is moot.\n"
f"{_r['text'][:800]}\n=== END LAST GEN ===\n")
# Final per-mode train-vs-deploy eval -- run for EVERY arm on the SAME fixed
# eval subset so the all-arms overlay reads identical numbers. For route/route2
# this is the absorption test: TRAIN keeps the quarantine knob on (still hacks),