-
Notifications
You must be signed in to change notification settings - Fork 659
fix(bridge): derive position_ids from attention_mask for left-padded input #1610
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Open
sohv
wants to merge
1
commit into
TransformerLensOrg:dev-4.x
Choose a base branch
from
sohv:fix/bridge-left-padding-positions
base: dev-4.x
Could not load branches
Branch not found: {{ refName }}
Loading
Could not load tags
Nothing to show
Loading
Are you sure you want to change the base?
Some commits from the old base branch may be removed from the timeline,
and old review comments may become outdated.
+199
−0
Open
Changes from all commits
Commits
File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
168 changes: 168 additions & 0 deletions
168
tests/integration/model_bridge/test_left_padding_positions.py
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,168 @@ | ||
| """Regression tests for left-padding position handling in TransformerBridge. | ||
|
|
||
| A causal LM's logits at a sequence's real token positions must not depend on how | ||
| that sequence is padded, provided the caller supplies the matching attention_mask. | ||
| Left padding shifts every real token's absolute position, so position_ids have to | ||
| be derived from the mask; without that the bridge silently returns wrong logits | ||
| and a wrong loss. See #1609. | ||
|
|
||
| Right padding is included as a control: causality already protects it, so it was | ||
| never affected and must stay that way. | ||
| """ | ||
|
|
||
| from __future__ import annotations | ||
|
|
||
| import pytest | ||
| import torch | ||
|
|
||
| from transformer_lens import utilities as utils | ||
|
|
||
| PAD_ID = 0 | ||
|
|
||
|
|
||
| def _pad(tokens: torch.Tensor, n_pad: int, side: str) -> tuple[torch.Tensor, torch.Tensor]: | ||
| """Pad `tokens` on `side`, returning (padded_tokens, attention_mask).""" | ||
| pads = torch.full((tokens.shape[0], n_pad), PAD_ID, dtype=tokens.dtype) | ||
| ones = torch.ones(tokens.shape, dtype=torch.long) | ||
| zeros = torch.zeros((tokens.shape[0], n_pad), dtype=torch.long) | ||
| if side == "left": | ||
| return torch.cat([pads, tokens], dim=1), torch.cat([zeros, ones], dim=1) | ||
| return torch.cat([tokens, pads], dim=1), torch.cat([ones, zeros], dim=1) | ||
|
|
||
|
|
||
| @pytest.fixture(scope="module") | ||
| def tokens(distilgpt2_bridge) -> torch.Tensor: | ||
| return distilgpt2_bridge.to_tokens("The capital of France is") | ||
|
|
||
|
|
||
| @pytest.mark.parametrize("side", ["left", "right"]) | ||
| @pytest.mark.parametrize("n_pad", [1, 3, 5]) | ||
| def test_logits_are_invariant_to_padding(distilgpt2_bridge, tokens, side, n_pad) -> None: | ||
| """Padding must not change the logits at a sequence's real positions.""" | ||
| padded, mask = _pad(tokens, n_pad, side) | ||
| real = slice(n_pad, None) if side == "left" else slice(None, tokens.shape[1]) | ||
|
|
||
| with torch.no_grad(): | ||
| baseline = distilgpt2_bridge(tokens, return_type="logits") | ||
| actual = distilgpt2_bridge(padded, attention_mask=mask, return_type="logits")[:, real] | ||
|
|
||
| assert torch.isfinite(actual).all() | ||
| torch.testing.assert_close(actual, baseline, rtol=1e-3, atol=1e-3) | ||
|
|
||
|
|
||
| @pytest.mark.parametrize("side", ["left", "right"]) | ||
| def test_compat_mode_logits_are_invariant_to_padding(distilgpt2_bridge_compat, side: str) -> None: | ||
| """enable_compatibility_mode() promises HookedTransformer-equivalent numerics, | ||
| which this property is part of.""" | ||
| tokens = distilgpt2_bridge_compat.to_tokens("The capital of France is") | ||
| n_pad = 3 | ||
| padded, mask = _pad(tokens, n_pad, side) | ||
| real = slice(n_pad, None) if side == "left" else slice(None, tokens.shape[1]) | ||
|
|
||
| with torch.no_grad(): | ||
| baseline = distilgpt2_bridge_compat(tokens, return_type="logits") | ||
| actual = distilgpt2_bridge_compat(padded, attention_mask=mask, return_type="logits")[ | ||
| :, real | ||
| ] | ||
|
|
||
| torch.testing.assert_close(actual, baseline, rtol=1e-3, atol=1e-3) | ||
|
|
||
|
|
||
| def test_derived_position_ids_match_hooked_transformer(distilgpt2_bridge, tokens) -> None: | ||
| """The derived positions must be the ones HookedTransformer would use, i.e. the | ||
| shared get_offset_position_ids helper rather than a parallel derivation.""" | ||
| n_pad = 3 | ||
| padded, mask = _pad(tokens, n_pad, "left") | ||
| expected = utils.get_offset_position_ids(0, mask) | ||
|
|
||
| with torch.no_grad(): | ||
| derived = distilgpt2_bridge(padded, attention_mask=mask, return_type="logits") | ||
| supplied = distilgpt2_bridge( | ||
| padded, attention_mask=mask, position_ids=expected, return_type="logits" | ||
| ) | ||
|
|
||
| torch.testing.assert_close(derived, supplied, rtol=1e-5, atol=1e-5) | ||
|
|
||
|
|
||
| def test_explicit_position_ids_take_precedence(distilgpt2_bridge, tokens) -> None: | ||
| """A caller-supplied position_ids must not be overwritten by the derivation.""" | ||
| n_pad = 3 | ||
| padded, mask = _pad(tokens, n_pad, "left") | ||
| derived_positions = utils.get_offset_position_ids(0, mask) | ||
| shifted = derived_positions + 1 # deliberately different, but still in range | ||
|
|
||
| with torch.no_grad(): | ||
| default = distilgpt2_bridge(padded, attention_mask=mask, return_type="logits") | ||
| overridden = distilgpt2_bridge( | ||
| padded, attention_mask=mask, position_ids=shifted, return_type="logits" | ||
| ) | ||
|
|
||
| assert not torch.allclose(default, overridden, rtol=1e-3, atol=1e-3) | ||
|
|
||
|
|
||
| @pytest.mark.parametrize("gap", [slice(3, 5), slice(1, 2)]) | ||
| def test_interior_mask_gap_uses_derived_positions(distilgpt2_bridge, tokens, gap) -> None: | ||
| """A mask gap that is not leading padding still shifts later positions, so the | ||
| derivation must fire for any mask, not only ones starting with a pad.""" | ||
| mask = torch.ones(tokens.shape, dtype=torch.long) | ||
| mask[0, gap] = 0 | ||
| gapped = tokens.clone() | ||
| gapped[0, gap] = PAD_ID | ||
| expected = utils.get_offset_position_ids(0, mask) | ||
|
|
||
| with torch.no_grad(): | ||
| derived = distilgpt2_bridge(gapped, attention_mask=mask, return_type="logits") | ||
| supplied = distilgpt2_bridge( | ||
| gapped, attention_mask=mask, position_ids=expected, return_type="logits" | ||
| ) | ||
|
|
||
| torch.testing.assert_close(derived, supplied, rtol=1e-5, atol=1e-5) | ||
|
|
||
|
|
||
| @pytest.mark.parametrize("mask_kind", ["all_ones", "right_padded"]) | ||
| def test_no_position_ids_injected_when_unnecessary(distilgpt2_bridge, tokens, mask_kind) -> None: | ||
| """Masks whose attended tokens already sit at their default positions must not | ||
| get position_ids injected: it is a no-op at best, and models whose forward does | ||
| not accept position_ids would raise. | ||
| """ | ||
| if mask_kind == "all_ones": | ||
| passed, mask = tokens, torch.ones(tokens.shape, dtype=torch.long) | ||
| else: | ||
| passed, mask = _pad(tokens, 3, "right") | ||
|
|
||
| seen: dict[str, object] = {} | ||
| original = distilgpt2_bridge.original_model.forward | ||
|
|
||
| def _spy(*args, **kwargs): | ||
| seen["position_ids"] = kwargs.get("position_ids") | ||
| return original(*args, **kwargs) | ||
|
|
||
| distilgpt2_bridge.original_model.forward = _spy | ||
| try: | ||
| with torch.no_grad(): | ||
| distilgpt2_bridge(passed, attention_mask=mask, return_type="logits") | ||
| finally: | ||
| distilgpt2_bridge.original_model.forward = original | ||
|
|
||
| assert seen["position_ids"] is None | ||
|
|
||
|
|
||
| def test_cached_step_with_left_padding(distilgpt2_bridge, tokens) -> None: | ||
| """With a KV cache the mask spans past+new while input_ids is only the new | ||
| token, so the derivation must be offset back to the tokens being passed.""" | ||
| n_pad = 3 | ||
| padded, mask = _pad(tokens, n_pad, "left") | ||
|
|
||
| with torch.no_grad(): | ||
| prefill = distilgpt2_bridge.original_model(padded, attention_mask=mask, use_cache=True) | ||
| new_token = tokens[:, -1:].clone() | ||
| extended = torch.cat([mask, torch.ones(1, 1, dtype=torch.long)], dim=1) | ||
| step = distilgpt2_bridge( | ||
| new_token, | ||
| attention_mask=extended, | ||
| past_key_values=prefill.past_key_values, | ||
| return_type="logits", | ||
| ) | ||
|
|
||
| assert step.shape[:2] == (1, 1) | ||
| assert torch.isfinite(step).all() |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -1551,6 +1551,37 @@ def forward( | |
| position_ids.masked_fill_(attention_mask == 0, 1) | ||
| kwargs["position_ids"] = position_ids | ||
|
|
||
| # Any masked-out token shifts the absolute position of every real token | ||
| # after it, so positions must be derived from the mask rather than left | ||
| # to HF's default arange. This is the same derivation HookedTransformer | ||
| # applies in pos_embed; without it the bridge silently returns wrong | ||
| # logits. An all-ones mask reduces to arange, so this is a no-op there. | ||
| # | ||
| # The mask spans any cached prefix as well as the new tokens, so it is | ||
| # offset back to just the tokens actually being passed — matching how | ||
| # AbstractAttention/PosEmbed use past_kv_pos_offset. | ||
| if ( | ||
| attention_mask is not None | ||
|
jlarson4 marked this conversation as resolved.
|
||
| and "position_ids" not in kwargs | ||
| and not _is_inputs_embeds | ||
| and attention_mask.ndim == 2 | ||
| and isinstance(input_ids, torch.Tensor) | ||
| and input_ids.ndim == 2 | ||
| and attention_mask.shape[1] >= input_ids.shape[1] | ||
| ): | ||
| _derived = utils.get_offset_position_ids(0, attention_mask) | ||
| _arange = torch.arange(attention_mask.shape[1], device=_derived.device) | ||
| # Only intervene when the mask actually moves an attended token off | ||
| # its default position — i.e. some masked token precedes a real one | ||
| # (left padding, or an interior gap). Pure right padding and an | ||
| # all-ones mask already agree with arange, and injecting position_ids | ||
| # there would be a no-op at best and unsupported by models whose | ||
| # forward does not take them at worst. | ||
| if bool(((_derived != _arange) & (attention_mask != 0)).any()): | ||
| kwargs["position_ids"] = _derived[ | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The predicate is a whole-batch |
||
| :, attention_mask.shape[1] - input_ids.shape[1] : | ||
| ] | ||
|
|
||
| if attention_mask is not None: | ||
| kwargs["attention_mask"] = attention_mask | ||
| if past_key_values is not None: | ||
|
|
||
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
For mRoPE architectures a left-padded batch now overrides the model's own 3-D derivation with a naive 2-D arange. HF only computes its rope index when
position_ids is None(modeling_qwen2_5_vl.py:1249) and silently expands 2-D input across all three streams (lines 822-823) and for fixed-signature remote-code models likeLLaDAModelLMthe same injection raisesTypeError, where base returned logits. Can injection be conditional on the target actually accepting and not already owning positions, the wayoutput_attentionsis guarded (bridge_core.py:1205-1218)?