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