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