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:
Nils
2026-07-18 00:16:08 +02:00
co-authored by Claude Fable 5
parent 3db6dd8fef
commit f0b7942c7a
5 changed files with 193 additions and 4 deletions
+35
View File
@@ -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 test. Verdict as registered: within +-5 -> no evidence that a native
attention read path breaches the consumption wall, with the init attention read path breaches the consumption wall, with the init
confound flagged as the one loose thread. 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.
+47 -2
View File
@@ -89,6 +89,42 @@ def splice_inner_iters(updates, inner_iters, inner_at, prompt_lens, dev, B):
return out 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, def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens,
k, use_checkpoint=False, feedforward=False, k, use_checkpoint=False, feedforward=False,
return_states=False, inner_iters=0, inner_at=None, 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() @torch.no_grad()
def generate_carry_c(looper, adapter, tok, input_ids, attention_mask, def generate_carry_c(looper, adapter, tok, input_ids, attention_mask,
k, p, max_new_tokens=10, feedforward=False, 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). """Greedy design-C generation (left-padded batch, uniform positions).
Appends p pause tokens, prefill-loops the prompt, carries through the 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), updates = [(torch.arange(B, device=dev),
torch.full((B,), n_prompt + j, device=dev, torch.full((B,), n_prompt + j, device=dev,
dtype=torch.long)) for j in range(p)] 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, anchor = torch.full((B,), n_prompt + p - 1, device=dev,
dtype=torch.long) dtype=torch.long)
inner = [(torch.arange(B, device=dev), anchor, True) inner = [(torch.arange(B, device=dev), anchor, True)
+27 -1
View File
@@ -43,6 +43,9 @@ def main():
ap.add_argument("--kvmem", default=None, metavar="KVMEM_PT", ap.add_argument("--kvmem", default=None, metavar="KVMEM_PT",
help="KVMemoryAdapter checkpoint: burst states become " help="KVMemoryAdapter checkpoint: burst states become "
"band-layer KV prefix entries at generation") "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() args = ap.parse_args()
model, tok = load_model(dtype=torch.bfloat16) model, tok = load_model(dtype=torch.bfloat16)
@@ -71,6 +74,29 @@ def main():
print(f"band-lora loaded: {args.bandlora} " print(f"band-lora loaded: {args.bandlora} "
f"(r={ck['rank']}, layers {ck['band'][0]}-{ck['band'][-1]})", f"(r={ck['rank']}, layers {ck['band'][0]}-{ck['band'][-1]})",
flush=True) 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( adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda() d=model.config.get_text_config().hidden_size).cuda()
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda")) adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
@@ -101,7 +127,7 @@ def main():
max_new_tokens=args.max_new, max_new_tokens=args.max_new,
feedforward=args.feedforward, feedforward=args.feedforward,
inner_iters=args.inner_iters, inner_iters=args.inner_iters,
kvmem=kvmem) kvmem=kvmem, symchain=symchain)
if kvmem is not None: if kvmem is not None:
from kv_memory import arm_memory from kv_memory import arm_memory
arm_memory(None) arm_memory(None)
+15
View File
@@ -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
+69 -1
View File
@@ -96,6 +96,12 @@ ap.add_argument("--kvmem", type=int, default=0, metavar="CODE",
ap.add_argument("--freeze-merge", action="store_true", ap.add_argument("--freeze-merge", action="store_true",
help="freeze the (warm-started) merge adapter; train only " help="freeze the (warm-started) merge adapter; train only "
"the kvmem adapter") "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", ap.add_argument("--teachstate", type=float, default=0.0, metavar="LAMBDA",
help="item 28 (Nils's variant): teacher-state distillation " help="item 28 (Nils's variant): teacher-state distillation "
"— frozen warm-start adapter runs the FULL cot (step " "— 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('.', '')}" f"_{str(ARGS.traj_fr).replace('.', '')}"
if (ARGS.traj_tf or ARGS.traj_fr) else "") + ( if (ARGS.traj_tf or ARGS.traj_fr) else "") + (
f"_kvm{ARGS.kvmem}" if ARGS.kvmem else "") + ( f"_kvm{ARGS.kvmem}" if ARGS.kvmem else "") + (
f"_sc{ARGS.symchain}" if ARGS.symchain else "") + (
"_fm" if ARGS.freeze_merge else "") + ( "_fm" if ARGS.freeze_merge else "") + (
f"_p{ARGS.base_pauses}" if ARGS.base_pauses >= 0 else "") + ( f"_p{ARGS.base_pauses}" if ARGS.base_pauses >= 0 else "") + (
f"_ln{ARGS.lensnoise.replace(',', '_')}" if ARGS.lensnoise else "") + ( f"_ln{ARGS.lensnoise.replace(',', '_')}" if ARGS.lensnoise else "") + (
@@ -197,6 +204,33 @@ def gen_staging_targets(tok, cot):
return out 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, def traj_burst_forward(looper, adapter, ids, msk, plens, m, T, tf_on, k,
mem_adapter=None): mem_adapter=None):
"""Item 29 forward: optional teacher-forced transition predictions, """Item 29 forward: optional teacher-forced transition predictions,
@@ -438,6 +472,18 @@ def main():
print(f"trajectory waypoints: {len(data)} items x {m} states " print(f"trajectory waypoints: {len(data)} items x {m} states "
f"({ARGS.traj_span} span, {time.time()-t0_:.0f}s; frozen " f"({ARGS.traj_span} span, {time.time()-t0_:.0f}s; frozen "
f"warm-start teacher, full cot)", flush=True) 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: if ARGS.teachstate:
assert ARGS.warm_start and ARGS.inner_iters, \ assert ARGS.warm_start and ARGS.inner_iters, \
"teachstate needs --warm-start (frozen teacher) + --inner-iters" "teachstate needs --warm-start (frozen teacher) + --inner-iters"
@@ -467,6 +513,9 @@ def main():
if mem_adapter is not None: if mem_adapter is not None:
groups.append({"params": list(mem_adapter.parameters()), "lr": LR, groups.append({"params": list(mem_adapter.parameters()), "lr": LR,
"base": LR}) "base": LR})
if sym_proj is not None:
groups.append({"params": list(sym_proj.parameters()), "lr": LR,
"base": LR})
if lora_params: if lora_params:
groups.append({"params": lora_params, "lr": ARGS.lora_lr, groups.append({"params": lora_params, "lr": ARGS.lora_lr,
"base": ARGS.lora_lr}) "base": ARGS.lora_lr})
@@ -498,7 +547,21 @@ def main():
for g in opt.param_groups: for g in opt.param_groups:
g["lr"] = g["base"] * lr_at(step) / LR g["lr"] = g["base"] * lr_at(step) / LR
ckw, itstates, tfp, frs = {}, None, None, None 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"]) T = (torch.stack([torch.as_tensor(b_["traj_states"])
for b_ in batch]).cuda() for b_ in batch]).cuda()
if (ARGS.traj_tf or ARGS.traj_fr) else None) 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] [p_ for p_ in adapter.parameters() if p_.requires_grad]
+ lora_params + lora_params
+ (list(mem_adapter.parameters()) if mem_adapter is not None + (list(mem_adapter.parameters()) if mem_adapter is not None
else [])
+ (list(sym_proj.parameters()) if sym_proj is not None
else []), 1.0) else []), 1.0)
opt.step() opt.step()
if ARGS.kvmem: if ARGS.kvmem:
@@ -638,6 +703,9 @@ def main():
if mem_adapter is not None: if mem_adapter is not None:
torch.save(mem_adapter.state_dict(), torch.save(mem_adapter.state_dict(),
OUT / f"kvmem_{TAG}_e{step+1}.pt") 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: if lora_params:
torch.save({"rank": ARGS.bandlora, "band": lora_band_layers, torch.save({"rank": ARGS.bandlora, "band": lora_band_layers,
"tensors": [p.detach().cpu() "tensors": [p.detach().cpu()