From f0b7942c7aecf2819400123a746bd973ee2bfcb8 Mon Sep 17 00:00:00 2001 From: Nils Date: Sat, 18 Jul 2026 00:16:08 +0200 Subject: [PATCH] =?UTF-8?q?item=2032=20pre-registered=20(Nils's=20synthesi?= =?UTF-8?q?s):=20discrete=20latent=20chain=20=E2=80=94=20lens-snapped=20to?= =?UTF-8?q?ken=20embeddings=20fed=20back=20via=20zero-init=20projector=20(?= =?UTF-8?q?ST=20top-32,=20TF/free-running=20arms,=20frozen=20arm-1=20merge?= =?UTF-8?q?);=20sym=5Fiterate=20in=20carry=5Fcommon,=20trainer/eval=20wiri?= =?UTF-8?q?ng?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Fable 5 --- results-loop/PROTOCOL_UNIFIED.md | 35 ++++++++++++++++ scripts/carry_common.py | 49 +++++++++++++++++++++- scripts/eval_carry_cot.py | 28 ++++++++++++- scripts/jobs/zzz_y_symchain.sh | 15 +++++++ scripts/train_carry_cot.py | 70 +++++++++++++++++++++++++++++++- 5 files changed, 193 insertions(+), 4 deletions(-) create mode 100644 scripts/jobs/zzz_y_symchain.sh diff --git a/results-loop/PROTOCOL_UNIFIED.md b/results-loop/PROTOCOL_UNIFIED.md index 34c5755..e3ef3c4 100644 --- a/results-loop/PROTOCOL_UNIFIED.md +++ b/results-loop/PROTOCOL_UNIFIED.md @@ -1038,3 +1038,38 @@ to explore any new channel; a headroom-bearing task would be a fairer test. Verdict as registered: within +-5 -> no evidence that a native attention read path breaches the consumption wall, with the init confound flagged as the one loose thread. + +32. **The discrete latent chain — "latent paper" (pre-registered + 2026-07-18 ~00:50, before running; Nils's synthesis: "the loop + needs paper, but we don't want that to be full tokens but still + latent space").** Diagnosis from the full matrix: every latent + medium lacked DISCRETENESS — tokens' magic is the snap + (error-correction per step), not visibility. The lens is a native + codebook: argmax over its readout quantizes any band state onto + the model's own symbol space. Design (each burst tick): + read s_{i-1} through the frozen lens; snap to a token + (straight-through over top-32, hard forward / soft gradient; tick + 0 = newline "a step begins"); feed E(token) back through a + ZERO-INIT projector alongside the analog carry: + x_i = merge(e, s_{i-1}) + proj(E(sym)). Two rails, matching the + microscopy's own two-channel algorithm: analog (plans/magnitudes) + + discrete (exact symbols). Writing and computing coincide by + construction: what the lens reads IS what gets transported — + item 25's lce loss (reused, λ=0.3) is now load-bearing, and its + proven concentration effect (10.3->1.9) doubles as the + quantization pressure that collapses the diffuse thinking-state + superposition (measured: pause states spread over ~hundreds of + tokens, committed states 1-2). Merge FROZEN at item-29 arm-1 + (39.1); only the 2.4M projector trains — the by-construction + ablation is 39.1. Arms: (a) sctf — teacher-forced symbols + (ground-truth deleted-step tokens; item-29's winning recipe); + (b) scst — free-running straight-through snaps. Eval always + hard-argmax free-running: 0:0 sanity + 2:0, n=256. Decision vs + 39.1: >44.1 = discreteness was the missing paper property, the + latent-chain program reopens (widenings pre-sketched: top-k + parallel snaps, L22/L26 depth rails, tape via slots/kvmem); + within +-5 = the discrete channel adds nothing over the analog + carry and the "loop = plan machine, tokens = executor" division + stands as final. Smokes: TF/ST train (loss arithmetic exact, + warm-start vals intact), eval generates with hard snaps. + Job: scripts/jobs/zzz_y_symchain.sh. diff --git a/scripts/carry_common.py b/scripts/carry_common.py index f70a911..916a9e9 100644 --- a/scripts/carry_common.py +++ b/scripts/carry_common.py @@ -89,6 +89,42 @@ def splice_inner_iters(updates, inner_iters, inner_at, prompt_lens, dev, B): return out +def sym_iterate(looper, adapter, proj, e, calls, S, X, rows, anchor, m, + lens_fn, embed_w, sym_tf=None, start_id=None, topk=32, + use_checkpoint=False, iter_states=None): + """Item 32: discrete latent chain at the anchor. Each tick reads the + previous anchor state through the lens, snaps it to a token + (straight-through over top-k) or takes the teacher token (sym_tf: + (B, m) ids, teacher forcing), and feeds that token's embedding back + through a zero-init projector ALONGSIDE the analog carry: + x_i = merge(e, s_{i-1}) + proj(E(sym)) + Tick 0 uses start_id (a newline: 'a step begins').""" + for i in range(m): + s_prev = S[rows, anchor] + if sym_tf is not None: + symb = embed_w[sym_tf[:, i]] + elif i == 0: + symb = embed_w[torch.full((rows.shape[0],), start_id, + device=e.device)] + else: + logits = lens_fn(s_prev).float() + p, idx = torch.softmax(logits, -1).topk(topk, dim=-1) + p = p / p.sum(-1, keepdim=True) + soft = (p.unsqueeze(-1) * embed_w[idx].float()).sum(-2) + hard = embed_w[idx[:, 0]].float() + symb = hard + soft - soft.detach() + x_new = (adapter(e[rows, anchor], s_prev).float() + + proj(symb.float())) + X = X.clone() + X[rows, anchor] = x_new.to(X.dtype) + S = (checkpoint(lambda X_: looper.band(X_, calls), X, + use_reentrant=False) if use_checkpoint + else looper.band(X, calls)) + if iter_states is not None: + iter_states.append(S[rows, anchor]) + return S, X + + def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens, k, use_checkpoint=False, feedforward=False, return_states=False, inner_iters=0, inner_at=None, @@ -128,7 +164,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, kvmem=None): + inner_iters=0, kvmem=None, symchain=None): """Greedy design-C generation (left-padded batch, uniform positions). Appends p pause tokens, prefill-loops the prompt, carries through the @@ -159,7 +195,16 @@ def generate_carry_c(looper, adapter, tok, input_ids, attention_mask, updates = [(torch.arange(B, device=dev), torch.full((B,), n_prompt + j, device=dev, dtype=torch.long)) for j in range(p)] - if inner_iters: + if symchain is not None and inner_iters: + anchor_sc = torch.full((B,), n_prompt + p - 1, device=dev, + dtype=torch.long) + S, X = sym_iterate( + looper, adapter, symchain["proj"], e, calls, S, X, + torch.arange(B, device=dev), anchor_sc, inner_iters, + symchain["lens_fn"], symchain["embed_w"], + start_id=symchain["start_id"]) + updates = updates # pauses (if any) already handled above + elif inner_iters: anchor = torch.full((B,), n_prompt + p - 1, device=dev, dtype=torch.long) inner = [(torch.arange(B, device=dev), anchor, True) diff --git a/scripts/eval_carry_cot.py b/scripts/eval_carry_cot.py index 4b37734..fc7122c 100644 --- a/scripts/eval_carry_cot.py +++ b/scripts/eval_carry_cot.py @@ -43,6 +43,9 @@ def main(): ap.add_argument("--kvmem", default=None, metavar="KVMEM_PT", help="KVMemoryAdapter checkpoint: burst states become " "band-layer KV prefix entries at generation") + ap.add_argument("--symchain", default=None, metavar="PROJ_PT", + help="item 32: discrete latent chain — zero-init " + "projector checkpoint; generation snaps hard argmax") args = ap.parse_args() model, tok = load_model(dtype=torch.bfloat16) @@ -71,6 +74,29 @@ def main(): print(f"band-lora loaded: {args.bandlora} " f"(r={ck['rank']}, layers {ck['band'][0]}-{ck['band'][-1]})", flush=True) + symchain = None + if args.symchain: + d_ = model.config.get_text_config().hidden_size + proj = torch.nn.Linear(d_, d_).cuda() + proj.load_state_dict(torch.load(args.symchain, map_location="cuda")) + proj.eval() + jbar = torch.load(Path(__file__).resolve().parent.parent + / "results/jbar.pt", map_location="cuda")["Jbar"] + from loop_common import BAND + J30 = jbar[BAND[1]].float() + tm = model.model.language_model + softcap = model.config.get_text_config().final_logit_softcapping + + def lens_fn(h): + x = tm.norm((h.float() @ J30.T).to(tm.norm.weight.dtype)) + lg = model.lm_head(x) + return softcap * torch.tanh(lg / softcap) if softcap else lg + + symchain = {"proj": proj, "lens_fn": lens_fn, + "embed_w": model.get_input_embeddings().weight.detach(), + "start_id": tok("\n", + add_special_tokens=False)["input_ids"][0]} + print(f"symchain loaded: {args.symchain}", flush=True) adapter = MergeAdapter( d=model.config.get_text_config().hidden_size).cuda() adapter.load_state_dict(torch.load(args.adapter, map_location="cuda")) @@ -101,7 +127,7 @@ def main(): max_new_tokens=args.max_new, feedforward=args.feedforward, inner_iters=args.inner_iters, - kvmem=kvmem) + kvmem=kvmem, symchain=symchain) if kvmem is not None: from kv_memory import arm_memory arm_memory(None) diff --git a/scripts/jobs/zzz_y_symchain.sh b/scripts/jobs/zzz_y_symchain.sh new file mode 100644 index 0000000..d5fdba7 --- /dev/null +++ b/scripts/jobs/zzz_y_symchain.sh @@ -0,0 +1,15 @@ +# 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_sc*.json results-loop/train_carrycot_b1_lt03_ii10_sc*_log.json results-loop/symproj_carrycot_b1_lt03_ii10_sc*_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 +for MODE in tf st; do + $P train_carry_cot.py --drop-steps 1 --pause-per-step 0 --base-pauses 0 \ + --inner-iters 10 --lensteach 0.3 --symchain $MODE --freeze-merge \ + --warm-start $A --steps 200 --lr 1e-3 + $P eval_carry_cot.py --adapter $A \ + --symchain $LOOP_OUT/symproj_carrycot_b1_lt03_ii10_sc${MODE}_fm_p0_e200.pt \ + --tag gsm_carrycot_b1_sc$MODE --grid 0:0,2:0 --n 256 --inner-iters 10 +done diff --git a/scripts/train_carry_cot.py b/scripts/train_carry_cot.py index c9a4663..5de8bee 100644 --- a/scripts/train_carry_cot.py +++ b/scripts/train_carry_cot.py @@ -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()