Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
168 changes: 168 additions & 0 deletions tests/integration/model_bridge/test_left_padding_positions.py
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()
31 changes: 31 additions & 0 deletions transformer_lens/model_bridge/transformer_bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (

Copy link
Copy Markdown
Collaborator

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 like LLaDAModelLM the same injection raises TypeError, where base returned logits. Can injection be conditional on the target actually accepting and not already owning positions, the way output_attentions is guarded (bridge_core.py:1205-1218)?

attention_mask is not None
Comment thread
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[

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The predicate is a whole-batch .any(), so a single left-padded row causes position_ids to be supplied for every row, including unpadded ones, which is what carries the mRoPE corruption into rows that had no padding at all. Is it possible to make the decision per-row?

:, 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:
Expand Down
Loading