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:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user