item 27 pre-registered: zero-pause internal band looping — single M=10 in-place burst at prompt end, iteration-aligned lens targets, ii0 inference ablation; carry_steps gains inplace updates, trainer gains --inner-iters/--base-pauses

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-17 00:21:20 +02:00
co-authored by Claude Fable 5
parent 0cffce876b
commit 44e814e78c
5 changed files with 149 additions and 16 deletions
+48 -7
View File
@@ -36,18 +36,26 @@ def prompt_prefill(looper, adapter, e, calls, prompt_mask, k):
def carry_steps(looper, adapter, e, calls, S, X, step_updates,
use_checkpoint=False):
"""Sequential scan. step_updates: list of (row_idx, pos) index tensors."""
for rows, pos in step_updates:
use_checkpoint=False, iter_states=None):
"""Sequential scan. step_updates: list of (row_idx, pos, [inplace])
index tensors. inplace=True re-iterates the SAME position, seeding
from its own previous band output (internal band looping) instead of
its left neighbor. iter_states: optional list — every inplace
update's fresh band output at pos is appended (for lens losses)."""
for upd in step_updates:
rows, pos = upd[0], upd[1]
inplace = len(upd) > 2 and upd[2]
if rows.numel() == 0:
continue
seed = S[rows, pos - 1]
seed = S[rows, pos] if inplace else S[rows, pos - 1]
x_new = adapter(e[rows, pos], seed)
X = X.clone()
X[rows, pos] = 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 inplace and iter_states is not None:
iter_states.append(S[rows, pos])
return S, X
@@ -63,9 +71,28 @@ def build_step_updates(prompt_lens, total_lens, device):
return updates
def splice_inner_iters(updates, inner_iters, inner_at, prompt_lens, dev, B):
"""Insert inner_iters in-place band iterations at the (batch-uniform
offset) anchor position, right after the scan first settles it."""
j_anchor = int((inner_at - prompt_lens).max())
rows = torch.arange(B, device=dev)
inner = [(rows, inner_at.to(dev), True)] * inner_iters
if j_anchor < 0:
# anchor is the last PROMPT position (no pauses at all):
# iterate there before the scan enters the visible tokens
return inner + updates
out = []
for j, u in enumerate(updates):
out.append(u)
if j == j_anchor:
out += inner
return out
def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens,
k, use_checkpoint=False, feedforward=False,
return_states=False):
return_states=False, inner_iters=0, inner_at=None,
iter_states=None):
"""Teacher-forced design-C forward (right-padded batch).
feedforward=True: pause-token control — same positions get the adapter as
@@ -87,15 +114,21 @@ def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens,
S, X = prompt_prefill(looper, adapter, e, calls, prompt_mask, k)
total_lens = attention_mask.sum(-1)
updates = build_step_updates(prompt_lens.to(dev), total_lens.to(dev), dev)
if inner_iters and inner_at is not None:
updates = splice_inner_iters(updates, inner_iters, inner_at,
prompt_lens.to(dev), dev,
input_ids.shape[0])
S, X = carry_steps(looper, adapter, e, calls, S, X, updates,
use_checkpoint=use_checkpoint)
use_checkpoint=use_checkpoint,
iter_states=iter_states)
out = looper.suffix_logits(S, calls)
return (out, S) if return_states else out
@torch.no_grad()
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):
"""Greedy design-C generation (left-padded batch, uniform positions).
Appends p pause tokens, prefill-loops the prompt, carries through the
@@ -126,6 +159,14 @@ 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:
anchor = torch.full((B,), n_prompt + p - 1, device=dev,
dtype=torch.long)
inner = [(torch.arange(B, device=dev), anchor, True)
for _ in range(inner_iters)]
# p=0: iterate at the last prompt position, before any
# visible token — no pause tokens involved
updates = (inner + updates) if p == 0 else (updates + inner)
S, X = carry_steps(looper, adapter, e, calls, S, X, updates)
else:
X = torch.cat([X_store, e[:, X_store.shape[1]:]], 1)