item 25 pre-registered: latent process supervision via differentiable lens readout — pause j trained to lens-encode deleted-step token j (λ=0.3/1.0 arms); carry_logits gains return_states

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-16 23:03:56 +02:00
co-authored by Claude Fable 5
parent f31e6e7e30
commit 64abfe66f3
4 changed files with 115 additions and 13 deletions
+6 -3
View File
@@ -64,7 +64,8 @@ def build_step_updates(prompt_lens, total_lens, device):
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):
"""Teacher-forced design-C forward (right-padded batch).
feedforward=True: pause-token control — same positions get the adapter as
@@ -81,13 +82,15 @@ def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens,
S = (checkpoint(lambda x_: looper.band(x_, calls), x,
use_reentrant=False) if use_checkpoint
else looper.band(x, calls))
return looper.suffix_logits(S, calls)
out = looper.suffix_logits(S, calls)
return (out, S) if return_states else out
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)
S, X = carry_steps(looper, adapter, e, calls, S, X, updates,
use_checkpoint=use_checkpoint)
return looper.suffix_logits(S, calls)
out = looper.suffix_logits(S, calls)
return (out, S) if return_states else out
@torch.no_grad()