item 31 pre-registered (Nils's design): synthetic memory tokens — burst states → per-band-layer KV prefix via KVMemoryAdapter (in-place attention wrap, bit-exact disarmed, gated silent init); composed with frozen item-29 arm-1; +29 tf+fr in-flight note (30.5, FR term hurts)

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-17 13:25:25 +02:00
co-authored by Claude Fable 5
parent 11604d0949
commit b18ffbd739
6 changed files with 231 additions and 11 deletions
+53 -8
View File
@@ -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()