item 32 pre-registered (Nils's synthesis): discrete latent chain — lens-snapped token embeddings fed back via zero-init projector (ST top-32, TF/free-running arms, frozen arm-1 merge); sym_iterate in carry_common, trainer/eval wiring
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -96,6 +96,12 @@ ap.add_argument("--kvmem", type=int, default=0, metavar="CODE",
|
||||
ap.add_argument("--freeze-merge", action="store_true",
|
||||
help="freeze the (warm-started) merge adapter; train only "
|
||||
"the kvmem adapter")
|
||||
ap.add_argument("--symchain", choices=("tf", "st"), default=None,
|
||||
help="item 32: discrete latent chain — each burst tick "
|
||||
"feeds back the lens-snapped token embedding through "
|
||||
"a zero-init projector. tf: teacher-forced symbols "
|
||||
"(ground-truth deleted-step tokens); st: free-running "
|
||||
"straight-through snaps")
|
||||
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 "
|
||||
@@ -117,6 +123,7 @@ TAG = ("carrycot_ff" if ARGS.feedforward else "carrycot") + (
|
||||
f"_{str(ARGS.traj_fr).replace('.', '')}"
|
||||
if (ARGS.traj_tf or ARGS.traj_fr) else "") + (
|
||||
f"_kvm{ARGS.kvmem}" if ARGS.kvmem else "") + (
|
||||
f"_sc{ARGS.symchain}" if ARGS.symchain 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 "") + (
|
||||
@@ -197,6 +204,33 @@ def gen_staging_targets(tok, cot):
|
||||
return out
|
||||
|
||||
|
||||
def sym_burst_forward(looper, adapter, proj, ids, msk, plens, m, sym_tf,
|
||||
k, lens_fn, embed_w, start_id):
|
||||
"""Item 32 forward: discrete latent chain at the anchor (teacher-forced
|
||||
or straight-through symbols), then the visible-token carry scan."""
|
||||
from carry_common import (prompt_prefill, build_step_updates,
|
||||
carry_steps, sym_iterate)
|
||||
dev = ids.device
|
||||
calls, _ = looper.capture(ids, msk, logits_to_keep=1)
|
||||
e = looper._hin[looper.l0].detach()
|
||||
B = ids.shape[0]
|
||||
ar = torch.arange(ids.shape[1], device=dev)
|
||||
pmask = ar[None, :] < plens[:, None].to(dev)
|
||||
S, X = prompt_prefill(looper, adapter, e, calls, pmask, k)
|
||||
rows = torch.arange(B, device=dev)
|
||||
anchor = (plens - 1).to(dev)
|
||||
states = []
|
||||
S, X = sym_iterate(looper, adapter, proj, e, calls, S, X, rows, anchor,
|
||||
m, lens_fn, embed_w, sym_tf=sym_tf,
|
||||
start_id=start_id, use_checkpoint=True,
|
||||
iter_states=states)
|
||||
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), states
|
||||
|
||||
|
||||
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,
|
||||
@@ -438,6 +472,18 @@ def main():
|
||||
print(f"trajectory waypoints: {len(data)} items x {m} states "
|
||||
f"({ARGS.traj_span} span, {time.time()-t0_:.0f}s; frozen "
|
||||
f"warm-start teacher, full cot)", flush=True)
|
||||
sym_proj, EMBED_W, START_ID = None, None, None
|
||||
if ARGS.symchain:
|
||||
assert ARGS.warm_start and ARGS.inner_iters and ARGS.lensteach, \
|
||||
"symchain needs --warm-start, --inner-iters, --lensteach"
|
||||
d_ = model.config.get_text_config().hidden_size
|
||||
sym_proj = torch.nn.Linear(d_, d_).cuda()
|
||||
torch.nn.init.zeros_(sym_proj.weight)
|
||||
torch.nn.init.zeros_(sym_proj.bias)
|
||||
EMBED_W = model.get_input_embeddings().weight.detach()
|
||||
START_ID = tok("\n", add_special_tokens=False)["input_ids"][0]
|
||||
print(f"symchain [{ARGS.symchain}]: zero-init proj "
|
||||
f"({d_}x{d_}), start_id={START_ID}", flush=True)
|
||||
if ARGS.teachstate:
|
||||
assert ARGS.warm_start and ARGS.inner_iters, \
|
||||
"teachstate needs --warm-start (frozen teacher) + --inner-iters"
|
||||
@@ -467,6 +513,9 @@ def main():
|
||||
if mem_adapter is not None:
|
||||
groups.append({"params": list(mem_adapter.parameters()), "lr": LR,
|
||||
"base": LR})
|
||||
if sym_proj is not None:
|
||||
groups.append({"params": list(sym_proj.parameters()), "lr": LR,
|
||||
"base": LR})
|
||||
if lora_params:
|
||||
groups.append({"params": lora_params, "lr": ARGS.lora_lr,
|
||||
"base": ARGS.lora_lr})
|
||||
@@ -498,7 +547,21 @@ 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 or ARGS.kvmem:
|
||||
if ARGS.symchain:
|
||||
sym_tf = None
|
||||
if ARGS.symchain == "tf":
|
||||
mat = torch.full((len(batch), ARGS.inner_iters), START_ID,
|
||||
dtype=torch.long)
|
||||
for b_i, it in enumerate(batch):
|
||||
tgt = it.get("lens_targets") or []
|
||||
row = [START_ID] + tgt[:-1]
|
||||
mat[b_i, :len(row)] = torch.tensor(row)
|
||||
sym_tf = mat.cuda()
|
||||
logits, itstates = sym_burst_forward(
|
||||
looper, adapter, sym_proj, ids, msk, plens,
|
||||
ARGS.inner_iters, sym_tf, K_PREFILL, lens_teach, EMBED_W,
|
||||
START_ID)
|
||||
elif 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)
|
||||
@@ -608,6 +671,8 @@ def main():
|
||||
[p_ for p_ in adapter.parameters() if p_.requires_grad]
|
||||
+ lora_params
|
||||
+ (list(mem_adapter.parameters()) if mem_adapter is not None
|
||||
else [])
|
||||
+ (list(sym_proj.parameters()) if sym_proj is not None
|
||||
else []), 1.0)
|
||||
opt.step()
|
||||
if ARGS.kvmem:
|
||||
@@ -638,6 +703,9 @@ def main():
|
||||
if mem_adapter is not None:
|
||||
torch.save(mem_adapter.state_dict(),
|
||||
OUT / f"kvmem_{TAG}_e{step+1}.pt")
|
||||
if sym_proj is not None:
|
||||
torch.save(sym_proj.state_dict(),
|
||||
OUT / f"symproj_{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