diff --git a/results-loop/PROTOCOL_UNIFIED.md b/results-loop/PROTOCOL_UNIFIED.md index 17586ed..d3b80a1 100644 --- a/results-loop/PROTOCOL_UNIFIED.md +++ b/results-loop/PROTOCOL_UNIFIED.md @@ -950,3 +950,35 @@ Nils's morning decision: clamp test vs pivot to the hybrid/A2 line matched accuracy; oracle meaningfully above 57.4 (early-stop rescues drift cases). Arm 2 (later): line-correctness / deferral head gating re-expansion. Job: scripts/jobs/zzz_w_metacog.sh. + +31. **Synthetic memory tokens: per-layer KV prefix (pre-registered + 2026-07-17 ~15:50, before running; Nils's design: "the one read + mechanism a frozen transformer natively possesses is attention over + the KV cache... adapter maps L30 state -> KV entries at band + layers"; injection variant chosen by Nils: per-layer KV prefix).** + The read-side attack, composed with item 29: warm-start the arm-1 + adapter (its burst COMPUTES — trajectory-trained, ltf 0.085 — but + is provably unread: ablation 8-8 p=1.0) and FREEZE it; train only + a KVMemoryAdapter (6.3M params: shared code-512 trunk + per-band- + layer k/v heads with per-layer geometry — gemma-4 mixes 256-d + sliding and 512-d global heads) mapping the 10 burst iterates to + post-RoPE (k,v) columns appended to every band layer's attention + during the answer scan and at generation. Engineering, validated + in smokes: the registered attention implementation is wrapped IN + PLACE (config name untouched -> mask construction identical; + disarmed = bit-exact, diff 0.0); zero-init value heads + learnable + per-layer gate bias (init -10) make the memory silent at init + (step-0 loss 0.108 = the arm-1 endpoint exactly); memory stays + armed through backward (checkpoint recompute). Output CE only, + 200 steps, lr 1e-3. Eval: 0:0 sanity + 2:0 matched (burst + + memory); the ablation is item 29's matched cell BY CONSTRUCTION + (same frozen adapter and burst, memory absent) = 39.1. Decision: + >44.1 (>=+5 over 39.1) = the consumption wall was a READ-PATH + problem and native attention over state-derived KV breaches it; + within +-5 = even natively readable computed states go unused -> + the wall is not about the read mechanism either, and the hybrid/ + metacog lines carry the program. In-flight note on item 29 arm 2 + (tf+fr): matched 30.5 — the free-running term HURT (vs 39.1 + tf-only), reinforcing the training-signal attribution: clean TF + transition gradients are the active ingredient. + Job: scripts/jobs/zzz_x_kvmem.sh. diff --git a/scripts/carry_common.py b/scripts/carry_common.py index a102167..f70a911 100644 --- a/scripts/carry_common.py +++ b/scripts/carry_common.py @@ -128,7 +128,7 @@ def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens, @torch.no_grad() def generate_carry_c(looper, adapter, tok, input_ids, attention_mask, k, p, max_new_tokens=10, feedforward=False, - inner_iters=0): + inner_iters=0, kvmem=None): """Greedy design-C generation (left-padded batch, uniform positions). Appends p pause tokens, prefill-loops the prompt, carries through the @@ -167,7 +167,12 @@ def generate_carry_c(looper, adapter, tok, input_ids, attention_mask, # p=0: iterate at the last prompt position, before any # visible token — no pause tokens involved updates = (inner + updates) if p == 0 else (updates + inner) - S, X = carry_steps(looper, adapter, e, calls, S, X, updates) + itst = [] if (kvmem is not None and inner_iters) else None + S, X = carry_steps(looper, adapter, e, calls, S, X, updates, + iter_states=itst) + if kvmem is not None and itst: + from kv_memory import arm_memory + arm_memory(kvmem(torch.stack(itst, 1))) else: X = torch.cat([X_store, e[:, X_store.shape[1]:]], 1) rows = torch.arange(B, device=dev) diff --git a/scripts/eval_carry_cot.py b/scripts/eval_carry_cot.py index 9cdbfe3..4b37734 100644 --- a/scripts/eval_carry_cot.py +++ b/scripts/eval_carry_cot.py @@ -40,11 +40,25 @@ def main(): help="load a lora_*_e*.pt loop-only band-LoRA checkpoint") ap.add_argument("--inner-iters", type=int, default=0, metavar="M", help="M in-place band iterations at the last pause") + ap.add_argument("--kvmem", default=None, metavar="KVMEM_PT", + help="KVMemoryAdapter checkpoint: burst states become " + "band-layer KV prefix entries at generation") args = ap.parse_args() model, tok = load_model(dtype=torch.bfloat16) tok.padding_side = "left" looper = BandLooper(model) + kvmem = None + if args.kvmem: + from kv_memory import install, KVMemoryAdapter + from loop_common import BAND + install(model) + sd = torch.load(args.kvmem, map_location="cuda") + code = sd["trunk.1.weight"].shape[0] + kvmem = KVMemoryAdapter(model, band=BAND, code=code).cuda() + kvmem.load_state_dict(sd) + kvmem.eval() + print(f"kv-memory loaded: {args.kvmem} (code={code})", flush=True) if args.bandlora: from lora_band import inject_band_lora ck = torch.load(args.bandlora, map_location="cuda") @@ -86,7 +100,11 @@ def main(): enc["attention_mask"], k, p, max_new_tokens=args.max_new, feedforward=args.feedforward, - inner_iters=args.inner_iters) + inner_iters=args.inner_iters, + kvmem=kvmem) + if kvmem is not None: + from kv_memory import arm_memory + arm_memory(None) for j, it in enumerate(chunk): txt = tok.decode(gen[j, enc["input_ids"].shape[1]:], skip_special_tokens=True) diff --git a/scripts/jobs/zzz_x_kvmem.sh b/scripts/jobs/zzz_x_kvmem.sh new file mode 100644 index 0000000..214bc71 --- /dev/null +++ b/scripts/jobs/zzz_x_kvmem.sh @@ -0,0 +1,13 @@ +# gpuq-in: results-loop/star_data.json results-loop/gsm_cot_data.json results-loop/adapter_carrycot_b1_ii10_tjs50_00_p0_e200.pt +# gpuq-out: results-loop/eval_gsm_carrycot_b1_kvm*.json results-loop/train_carrycot_b1_ii10_kvm512_fm_p0_log.json results-loop/kvmem_carrycot_b1_ii10_kvm512_fm_p0_e200.pt +git pull origin main -q 2>/dev/null +P=/home/nils/jspace/.venv/bin/python +export JLENS_MODEL=google/gemma-4-E2B-it LOOP_OUT=/home/nils/jspace/results-loop +cd /home/nils/jspace/scripts +A=$LOOP_OUT/adapter_carrycot_b1_ii10_tjs50_00_p0_e200.pt +$P train_carry_cot.py --drop-steps 1 --pause-per-step 0 --base-pauses 0 \ + --inner-iters 10 --kvmem 512 --freeze-merge \ + --warm-start $A --steps 200 --lr 1e-3 +$P eval_carry_cot.py --adapter $A \ + --kvmem $LOOP_OUT/kvmem_carrycot_b1_ii10_kvm512_fm_p0_e200.pt \ + --tag gsm_carrycot_b1_kvm512 --grid 0:0,2:0 --n 256 --inner-iters 10 diff --git a/scripts/kv_memory.py b/scripts/kv_memory.py new file mode 100644 index 0000000..dca4e95 --- /dev/null +++ b/scripts/kv_memory.py @@ -0,0 +1,107 @@ +"""Item 31: synthetic memory tokens — per-layer KV prefix (Nils's design). + +"The one read mechanism a frozen transformer natively possesses is +attention over the KV cache." A KVMemoryAdapter maps carried L30 states +to (k, v) entries appended to every band layer's attention — frozen +heads consume the state through their existing machinery. Keys live in +post-RoPE space (prefix-tuning convention: learned keys place themselves +where frozen queries already point). + +Activation is global-flag-scoped like lora_band.LoopLoRA: a custom +attention interface "kvmem" wraps the stock eager path; when _MEM is +armed and the module is a band layer, memory kv columns are appended +(mask columns visible to every query). +""" +import torch +import torch.nn as nn + +from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS +from transformers.models.gemma4.modeling_gemma4 import ( + eager_attention_forward) + +_MEM = {"on": False, "kv": None, "band": (14, 30), "base": None} + + +def arm_memory(kv): + """kv: dict layer_idx -> (k, v), each (B, kv_heads, M, head_dim).""" + _MEM["kv"] = kv + _MEM["on"] = kv is not None + + +def kvmem_attention(module, query, key, value, attention_mask, **kwargs): + if (_MEM["on"] and _MEM["kv"] is not None + and getattr(module, "layer_idx", -1) in _MEM["kv"]): + mk, mv, gate = _MEM["kv"][module.layer_idx] + m = mk.shape[2] + key = torch.cat([key, mk.to(dtype=key.dtype, device=key.device)], 2) + value = torch.cat([value, mv.to(dtype=value.dtype, + device=value.device)], 2) + if attention_mask is not None: + pad = torch.zeros(*attention_mask.shape[:-1], m, + dtype=attention_mask.dtype, + device=attention_mask.device) + gate.to( + attention_mask.dtype) + attention_mask = torch.cat([attention_mask, pad], -1) + base = _MEM["base"] or eager_attention_forward + return base(module, query, key, value, attention_mask, **kwargs) + + +def install(model): + """Wrap the registered implementation IN PLACE — the config name is + untouched, so mask construction and every other branch stays + bit-identical; only the (q,k,v,mask) seam gains the memory append.""" + if _MEM.get("installed"): + return + impl = model.config._attn_implementation + _MEM["base"] = ALL_ATTENTION_FUNCTIONS.get_interface( + impl, eager_attention_forward) + try: + ALL_ATTENTION_FUNCTIONS[impl] = kvmem_attention + except TypeError: + ALL_ATTENTION_FUNCTIONS.register(impl, kvmem_attention) + _MEM["installed"] = True + + +class KVMemoryAdapter(nn.Module): + """Carried states (B, M, d) -> per-band-layer post-RoPE (k, v). + Per-layer geometry (head_dim, kv heads) read off the live modules — + gemma-4 mixes sliding (256) and global (512) head dims in the band.""" + + def __init__(self, model, band=(14, 30), code=512): + super().__init__() + self.band = band + tm = model.model.language_model + d = model.config.get_text_config().hidden_size + self.trunk = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, code), + nn.GELU()) + self.geom = {} + self.k_heads = nn.ModuleDict() + self.v_heads = nn.ModuleDict() + tc = model.config.get_text_config() + for l in range(band[0], band[1] + 1): + att = tm.layers[l].self_attn + hd = att.head_dim + kvh = (tc.num_global_key_value_heads + if getattr(att, "use_alternative_attention", False) + else tc.num_key_value_heads) + self.geom[l] = (kvh, hd) + self.k_heads[str(l)] = nn.Linear(code, kvh * hd) + self.v_heads[str(l)] = nn.Linear(code, kvh * hd) + nn.init.zeros_(self.v_heads[str(l)].weight) + nn.init.zeros_(self.v_heads[str(l)].bias) + # learnable per-layer attention gate on the memory columns: + # init -10 -> memory is invisible until training opens it + self.gate = nn.ParameterDict( + {str(l): nn.Parameter(torch.tensor(-10.0)) + for l in range(band[0], band[1] + 1)}) + + def forward(self, states): + B, M, _ = states.shape + c = self.trunk(states.float()) + kv = {} + for l in range(self.band[0], self.band[1] + 1): + kvh, hd = self.geom[l] + k = self.k_heads[str(l)](c).view(B, M, kvh, hd).transpose(1, 2) + v = self.v_heads[str(l)](c).view(B, M, kvh, hd).transpose(1, 2) + kv[l] = (k, v, self.gate[str(l)]) + return kv diff --git a/scripts/train_carry_cot.py b/scripts/train_carry_cot.py index e01688f..c9a4663 100644 --- a/scripts/train_carry_cot.py +++ b/scripts/train_carry_cot.py @@ -89,6 +89,13 @@ ap.add_argument("--traj-fr", type=float, default=0.0, metavar="LAMBDA", ap.add_argument("--traj-span", choices=("step", "full"), default="step", help="waypoints sampled evenly across the deleted step " "(job 1, d=1) or the full CoT (job 2, answer-only)") +ap.add_argument("--kvmem", type=int, default=0, metavar="CODE", + help="item 31: KVMemoryAdapter (code width) — burst states " + "become per-band-layer KV prefix entries readable by " + "frozen attention during the answer scan") +ap.add_argument("--freeze-merge", action="store_true", + help="freeze the (warm-started) merge adapter; train only " + "the kvmem adapter") ap.add_argument("--teachstate", type=float, default=0.0, metavar="LAMBDA", help="item 28 (Nils's variant): teacher-state distillation " "— frozen warm-start adapter runs the FULL cot (step " @@ -109,6 +116,8 @@ TAG = ("carrycot_ff" if ARGS.feedforward else "carrycot") + ( f"_tj{ARGS.traj_span[0]}{str(ARGS.traj_tf).replace('.', '')}" f"_{str(ARGS.traj_fr).replace('.', '')}" if (ARGS.traj_tf or ARGS.traj_fr) else "") + ( + f"_kvm{ARGS.kvmem}" if ARGS.kvmem else "") + ( + "_fm" if ARGS.freeze_merge else "") + ( f"_p{ARGS.base_pauses}" if ARGS.base_pauses >= 0 else "") + ( f"_ln{ARGS.lensnoise.replace(',', '_')}" if ARGS.lensnoise else "") + ( f"_s{ARGS.seed}" if ARGS.seed else "") + ARGS.tag_suffix @@ -188,7 +197,8 @@ def gen_staging_targets(tok, cot): return out -def traj_burst_forward(looper, adapter, ids, msk, plens, m, T, tf_on, k): +def traj_burst_forward(looper, adapter, ids, msk, plens, m, T, tf_on, k, + mem_adapter=None): """Item 29 forward: optional teacher-forced transition predictions, then the free-running burst (whose final state seeds the answer scan, matching inference), then the visible-token carry. Returns @@ -224,11 +234,17 @@ def traj_burst_forward(looper, adapter, ids, msk, plens, m, T, tf_on, k): seed = s0 if i == 0 else fr_states[-1] S, X = upd(S, X, seed) fr_states.append(S[rows, anchor]) + if mem_adapter is not None: + from kv_memory import arm_memory + arm_memory(mem_adapter(torch.stack(fr_states, 1))) total = msk.sum(-1) updates = build_step_updates(plens.to(dev), total.to(dev), dev) S, X = carry_steps(looper, adapter, e, calls, S, X, updates, use_checkpoint=True) - return looper.suffix_logits(S, calls), tf_preds, fr_states + out = looper.suffix_logits(S, calls) + # NOTE: memory stays armed through backward (checkpoint recompute + # must see the same graph); the caller disarms after opt.step(). + return out, tf_preds, fr_states def lr_at(step): @@ -339,6 +355,20 @@ def main(): f"{sum(p.numel() for p in lora_params)/1e6:.1f}M params, " f"layers {lora_band_layers[0]}-{lora_band_layers[-1]}, " f"lr={ARGS.lora_lr}", flush=True) + mem_adapter = None + if ARGS.kvmem: + from kv_memory import install, KVMemoryAdapter + from loop_common import BAND + install(model) + mem_adapter = KVMemoryAdapter(model, band=BAND, + code=ARGS.kvmem).cuda() + print(f"kv-memory adapter: code={ARGS.kvmem}, " + f"{sum(p.numel() for p in mem_adapter.parameters())/1e6:.1f}M " + f"params, band {BAND}", flush=True) + if ARGS.freeze_merge: + for p_ in adapter.parameters(): + p_.requires_grad_(False) + print("merge adapter FROZEN", flush=True) lens_teach = None if ARGS.lensteach or ARGS.lensteach_gen: from loop_common import BAND @@ -431,7 +461,12 @@ def main(): print(f"teacher states: {len(todo)} captured ({time.time()-t0_:.0f}s;" f" frozen warm-start adapter, full cot, band-exit at the " f"deleted step's last token)", flush=True) - groups = [{"params": list(adapter.parameters()), "lr": LR, "base": LR}] + groups = ([] if ARGS.freeze_merge else + [{"params": list(adapter.parameters()), "lr": LR, + "base": LR}]) + if mem_adapter is not None: + groups.append({"params": list(mem_adapter.parameters()), "lr": LR, + "base": LR}) if lora_params: groups.append({"params": lora_params, "lr": ARGS.lora_lr, "base": ARGS.lora_lr}) @@ -463,12 +498,13 @@ def main(): for g in opt.param_groups: g["lr"] = g["base"] * lr_at(step) / LR ckw, itstates, tfp, frs = {}, None, None, None - if ARGS.traj_tf or ARGS.traj_fr: - T = torch.stack([torch.as_tensor(b_["traj_states"]) - for b_ in batch]).cuda() + if ARGS.traj_tf or ARGS.traj_fr or ARGS.kvmem: + T = (torch.stack([torch.as_tensor(b_["traj_states"]) + for b_ in batch]).cuda() + if (ARGS.traj_tf or ARGS.traj_fr) else None) logits, tfp, frs = traj_burst_forward( looper, adapter, ids, msk, plens, ARGS.inner_iters, T, - bool(ARGS.traj_tf), K_PREFILL) + bool(ARGS.traj_tf), K_PREFILL, mem_adapter=mem_adapter) else: if ARGS.inner_iters: extras = torch.tensor([b_.get("extra_pauses", 0) @@ -569,8 +605,14 @@ def main(): opt.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_( - list(adapter.parameters()) + lora_params, 1.0) + [p_ for p_ in adapter.parameters() if p_.requires_grad] + + lora_params + + (list(mem_adapter.parameters()) if mem_adapter is not None + else []), 1.0) opt.step() + if ARGS.kvmem: + from kv_memory import arm_memory + arm_memory(None) log.append({"step": step, "loss": loss.item(), "lce": lce_val, "lgen": lgen_val, "lts": lts_val, "ltf": ltf_val, "lfr": lfr_val}) @@ -593,6 +635,9 @@ def main(): adapter.noise_on = True sd = (adapter.base if ARGS.lensnoise else adapter).state_dict() torch.save(sd, OUT / f"adapter_{TAG}_e{step+1}.pt") + if mem_adapter is not None: + torch.save(mem_adapter.state_dict(), + OUT / f"kvmem_{TAG}_e{step+1}.pt") if lora_params: torch.save({"rank": ARGS.bandlora, "band": lora_band_layers, "tensors": [p.detach().cpu()