fix(bridge): gate batched-list position_ids on the target model - #1627
fix(bridge): gate batched-list position_ids on the target model#1627sohv wants to merge 1 commit into
Conversation
Batched list input builds an attention_mask and position_ids itself so pad tokens don't contaminate the forward. The mask is safe for any model, but the position_ids were handed over unchecked: a forward taking neither position_ids nor **kwargs raises TypeError where it would have returned logits. This is the gap jlarson4 raised while reviewing TransformerLensOrg#1610. That PR added _accepts_derived_position_ids() and gated the main forward() derivation, but these two sites were left for a follow-up because no model could be shown to fail there. The LLaDA test harness builds a fixed-signature forward in process, which reproduces it: TypeError: TinyLLaDAModelLM.forward() got an unexpected keyword argument 'position_ids' Gate both sites on the same helper. The attention_mask stays unconditional -- it is safe everywhere, and withholding it would reintroduce the padding contamination this branch exists to prevent. A single unbatched string was never affected, and the test asserts that alongside the batched case. The regression test wraps its forward spy in functools.wraps: the gate reads that forward's signature, so a bare (*args, **kwargs) wrapper would look like it accepts position_ids and silently defeat the check under test. Fixes TransformerLensOrg#1626 Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
| forward_kwargs["position_ids"] = position_ids | ||
| # Same target gate as the forward() path: the mask is safe | ||
| # for every model, the derived positions are not (#1626). | ||
| if self._accepts_derived_position_ids(): |
There was a problem hiding this comment.
For a gate=False model that accepts position_ids (opt-125m), this gate makes cached steps fall into the torch.full(total_len - 1) fallback at line 2220, which counts pad slots, causing logit drift. Can you gate the continuation's derive/fallback injection starting at line 2213-2226 on the same helper?
| torch.testing.assert_close(bridge_logits, reference_logits, rtol=1e-5, atol=1e-6) | ||
|
|
||
|
|
||
| def test_batched_list_input_does_not_inject_unsupported_position_ids() -> None: |
There was a problem hiding this comment.
The _generate_tokens portion of this fix is untested here and neither fixed-signature architecture can reach that site. Can you add a cached-vs-uncached batched-generate parity test on a gate=False model?
|
@jlarson4 A quick request to review this PR and merge it if you have no issues. I reviewed the PR again today but let me know if there is anything to be addressed here. |
The review just posted! Sorry I didn't get it to you yesterday, was down with the flu. |
|
@jlarson4 It's totally alright and thank you for highlighting these issues. I will work on fixing them and will update the PR soon. |
Fixes #1626.
Batched list input builds an
attention_maskandposition_idsitself so pad tokens don't contaminate the forward. The mask is safe for any model, but theposition_idswere handed over unchecked — a forward taking neitherposition_idsnor**kwargsraisesTypeErrorwhere it would otherwise have returned logits.This is the gap that was raised while reviewing #1610. That PR added
_accepts_derived_position_ids()and gated the mainforward()derivation, but we left these two sites for a follow-up because I couldn't produce a model that demonstrably failed there. The LLaDA test harness builds a fixed-signature forward in-process, which reproduces it:What I changed
Both sites now call the same
_accepts_derived_position_ids()helper, so there's no new predicate. Theattention_maskstays unconditional -it's safe for every model, and withholding it would reintroduce the padding contamination the branch exists to prevent.Verification
One regression test in the LLaDA suite, red on
dev-4.xwith theTypeErrorabove and green with the fix. It asserts both halves: noposition_idsreaches the model, and theattention_maskstill does. A single unbatched string was never affected and is covered as a control.Full unit + integration + acceptance: 6594 passed, 1 failed. The failure is
test_bridge_hooked_parity_multi_step_optimization, which fails identically on unmodified code and is--ignored on the macOS CI job. mypy clean.