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
+47 -2
View File
@@ -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)