diff --git a/tests/unit/model_bridge/test_loss_attention_mask.py b/tests/unit/model_bridge/test_loss_attention_mask.py new file mode 100644 index 000000000..548bff75b --- /dev/null +++ b/tests/unit/model_bridge/test_loss_attention_mask.py @@ -0,0 +1,127 @@ +"""Regression tests for padding-aware TransformerBridge causal loss.""" + +from __future__ import annotations + +import pytest +import torch + +from transformer_lens.config import TransformerBridgeConfig +from transformer_lens.model_bridge import TransformerBridge + + +def _bridge() -> TransformerBridge: + cfg = TransformerBridgeConfig( + d_model=32, + d_head=8, + n_heads=4, + n_layers=2, + n_ctx=6, + d_vocab=32, + d_mlp=64, + act_fn="gelu", + normalization_type="LN", + seed=7, + initializer_range=0.2, + ) + return TransformerBridge.boot_native(cfg) + + +def _extract_loss(output: torch.Tensor | tuple[torch.Tensor, torch.Tensor]) -> torch.Tensor: + return output[1] if isinstance(output, tuple) else output + + +@pytest.mark.parametrize("return_type", ["loss", "both"]) +def test_forward_loss_ignores_masked_padding_tokens(return_type: str) -> None: + bridge = _bridge() + attention_mask = torch.tensor( + [ + [1, 1, 1, 0, 0, 0], + [1, 1, 1, 1, 1, 1], + ] + ) + token_batches = ( + torch.tensor( + [ + [1, 2, 3, 0, 0, 0], + [4, 5, 6, 7, 8, 9], + ] + ), + torch.tensor( + [ + [1, 2, 3, 31, 30, 29], + [4, 5, 6, 7, 8, 9], + ] + ), + ) + + losses = [] + for tokens in token_batches: + output = bridge(tokens, attention_mask=attention_mask, return_type=return_type) + loss = _extract_loss(output) + logits = bridge(tokens, attention_mask=attention_mask, return_type="logits") + expected = bridge.loss_fn(logits, tokens, attention_mask=attention_mask) + + torch.testing.assert_close(loss, expected) + losses.append(loss) + + torch.testing.assert_close(losses[0], losses[1]) + + +def test_forward_loss_per_token_zeros_masked_transitions() -> None: + bridge = _bridge() + tokens = torch.tensor( + [ + [1, 2, 3, 0, 0, 0], + [4, 5, 6, 7, 8, 9], + ] + ) + attention_mask = torch.tensor( + [ + [1, 1, 1, 0, 0, 0], + [1, 1, 1, 1, 1, 1], + ] + ) + + loss = bridge( + tokens, + attention_mask=attention_mask, + return_type="loss", + loss_per_token=True, + ) + next_token_mask = torch.logical_and(attention_mask[:, :-1], attention_mask[:, 1:]) + + assert torch.count_nonzero(loss[~next_token_mask]) == 0 + + +def test_forward_loss_is_finite_with_left_padding() -> None: + bridge = _bridge() + tokens = torch.tensor( + [ + [0, 0, 0, 1, 2, 3], + [4, 5, 6, 7, 8, 9], + ] + ) + attention_mask = torch.tensor( + [ + [0, 0, 0, 1, 1, 1], + [1, 1, 1, 1, 1, 1], + ] + ) + position_ids = attention_mask.long().cumsum(-1) - 1 + position_ids.masked_fill_(attention_mask == 0, 1) + + logits = bridge( + tokens, + attention_mask=attention_mask, + position_ids=position_ids, + return_type="logits", + ) + loss = bridge( + tokens, + attention_mask=attention_mask, + position_ids=position_ids, + return_type="loss", + ) + + assert torch.isfinite(logits).all() + assert torch.isfinite(loss) diff --git a/transformer_lens/model_bridge/bridge_core.py b/transformer_lens/model_bridge/bridge_core.py index f361478ce..9f7da93ed 100644 --- a/transformer_lens/model_bridge/bridge_core.py +++ b/transformer_lens/model_bridge/bridge_core.py @@ -345,6 +345,8 @@ def loss_fn( """Cross-entropy loss matching HookedTransformer's formula (log_softmax + gather).""" if tokens.device != logits.device: tokens = tokens.to(logits.device) + if attention_mask is not None and attention_mask.device != logits.device: + attention_mask = attention_mask.to(logits.device) return lm_cross_entropy_loss(logits, tokens, attention_mask, per_token) def _finalize_return( @@ -353,6 +355,7 @@ def _finalize_return( logits: Optional[torch.Tensor], input_ids: Optional[torch.Tensor], *, + attention_mask: Optional[torch.Tensor] = None, is_audio_model: bool = False, is_visual_model: bool = False, inputs_embeds_was_used: bool = False, @@ -384,7 +387,12 @@ def _finalize_return( ) assert isinstance(logits, torch.Tensor), f"Expected logits tensor, got {type(logits)}" assert input_ids is not None, "input_ids required for return_type='loss'" - return self.loss_fn(logits, input_ids, per_token=loss_per_token) + return self.loss_fn( + logits, + input_ids, + attention_mask=attention_mask, + per_token=loss_per_token, + ) if return_type == "both": if is_audio_model: raise ValueError( @@ -404,7 +412,12 @@ def _finalize_return( ) assert isinstance(logits, torch.Tensor), f"Expected logits tensor, got {type(logits)}" assert input_ids is not None, "input_ids required for return_type='both'" - loss = self.loss_fn(logits, input_ids, per_token=loss_per_token) + loss = self.loss_fn( + logits, + input_ids, + attention_mask=attention_mask, + per_token=loss_per_token, + ) return (logits, loss) if return_type == "predictions": assert self.tokenizer is not None, "Tokenizer required for return_type='predictions'" diff --git a/transformer_lens/model_bridge/sources/native/model.py b/transformer_lens/model_bridge/sources/native/model.py index 5da5a1b6e..02a3aa9fb 100644 --- a/transformer_lens/model_bridge/sources/native/model.py +++ b/transformer_lens/model_bridge/sources/native/model.py @@ -329,6 +329,9 @@ def forward( scores = scores.masked_fill(block_mask, float("-inf")) pattern = F.softmax(scores, dim=-1) + # Fully masked padding queries softmax to NaN; overwrite masked entries + # so those rows contribute a zero attention update instead of poisoning later layers. + pattern = pattern.masked_fill(block_mask, 0.0) attn = torch.matmul(pattern, v).transpose(1, 2).contiguous().view(batch, seq, -1) out = self.o(attn) diff --git a/transformer_lens/model_bridge/transformer_bridge.py b/transformer_lens/model_bridge/transformer_bridge.py index c0318a5b6..2c376dae2 100644 --- a/transformer_lens/model_bridge/transformer_bridge.py +++ b/transformer_lens/model_bridge/transformer_bridge.py @@ -1658,6 +1658,7 @@ def forward( return_type, logits, input_ids, + attention_mask=attention_mask, is_audio_model=getattr(self.cfg, "is_audio_model", False), is_visual_model=getattr(self.cfg, "is_visual_model", False), inputs_embeds_was_used=_is_inputs_embeds, diff --git a/transformer_lens/utilities/lm_utils.py b/transformer_lens/utilities/lm_utils.py index a3d4f7932..db83c7ad3 100644 --- a/transformer_lens/utilities/lm_utils.py +++ b/transformer_lens/utilities/lm_utils.py @@ -17,7 +17,7 @@ def lm_cross_entropy_loss( tokens: Int[torch.Tensor, "batch pos"], attention_mask: Optional[Int[torch.Tensor, "batch pos"]] = None, per_token: bool = False, -) -> Union[Float[torch.Tensor, ""], Float[torch.Tensor, "batch pos"]]: +) -> Union[Float[torch.Tensor, ""], Float[torch.Tensor, "batch pos_minus_one"]]: """Cross entropy loss for the language model, gives the loss for predicting the NEXT token. Args: @@ -37,7 +37,7 @@ def lm_cross_entropy_loss( # Ignore token positions which are masked out or where the next token is masked out # (generally padding tokens) next_token_mask = torch.logical_and(attention_mask[:, :-1], attention_mask[:, 1:]) - predicted_log_probs *= next_token_mask + predicted_log_probs = predicted_log_probs.masked_fill(~next_token_mask, 0.0) n_tokens = next_token_mask.sum().item() else: n_tokens = predicted_log_probs.numel()