Skip to content

Fix masked causal loss in TransformerBridge - #1608

Merged
jlarson4 merged 2 commits into
TransformerLensOrg:dev-4.xfrom
emerardd:fix/bridge-loss-attention-mask
Aug 6, 2026
Merged

Fix masked causal loss in TransformerBridge#1608
jlarson4 merged 2 commits into
TransformerLensOrg:dev-4.xfrom
emerardd:fix/bridge-loss-attention-mask

Conversation

@emerardd

@emerardd emerardd commented Aug 5, 2026

Copy link
Copy Markdown
Contributor

Description

Propagate the forward-pass attention_mask into TransformerBridge causal loss so padding transitions do not affect return_type="loss", return_type="both", or per-token loss.

The fix also:

  • moves the loss mask to the logits device when needed;
  • uses NaN-safe masking for ignored token losses;
  • zeros fully masked Native attention pattern entries so left padding cannot poison later layers; and
  • corrects the per-token loss return annotation to reflect its seq_len - 1 shape.

Fixes #1607

Type of change

  • Bug fix (non-breaking change which fixes an issue)

Screenshots

Not applicable.

Validation

  • Regression tests: 4 passed
  • Affected Bridge unit tests: 102 passed, 1 skipped
  • Changed files: pycln, isort, and Black clean
  • git diff --check: clean
  • uv run mypy .: success on 422 source files

Checklist

  • I have commented my code, particularly in hard-to-understand areas
  • Documentation is not required for this internal correctness fix
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective
  • The complete unit-test suite passes locally; only the affected unit-test surface was run
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

@emerardd
emerardd marked this pull request as ready for review August 5, 2026 07:50
sohv added a commit to sohv/TransformerLens that referenced this pull request Aug 5, 2026
…input

TransformerBridge.forward() did not derive position_ids from a supplied
attention_mask, so left-padded input silently got the wrong absolute positions —
no error, no NaN, just wrong logits and a wrong loss. On gpt2 the loss for one
prompt moved from 4.503170 unpadded to 11.154946 with three left pads, while
HookedTransformer stays invariant (drift ~1e-06). Right padding was never
affected, since causality already protects it.

transformer_bridge.py derived position_ids only for batched *list* input, so
pre-tokenized tensors fell through to HF's plain arange and the padding offset
was never removed. This extends the same correction the list-input branch
already applies: when a mask is supplied, position_ids are absent, and the mask
indicates left padding, positions are computed from the mask. An explicitly
supplied position_ids still wins.

The bridge was also inconsistent with itself before this — the same batch gave
different logits depending on whether it was passed as strings or token IDs
(max |logit diff| 4.142e+01) — and enable_compatibility_mode(), which documents
"HookedTransformer-equivalent numerics", diverged on left-padded input while
matching exactly on unpadded input.

Adds integration regression tests covering logit invariance under both padding
sides, the same property in compatibility mode, and agreement between the
derived and an explicitly supplied position_ids. The right-padding cases are
controls: they pass with and without the fix. The tests live in the integration
tier because left padding makes a fully masked query row, which the Native
attention path turns into NaN until the masked-softmax fix in TransformerLensOrg#1608 lands.

Fixes TransformerLensOrg#1609.

Co-Authored-By: Claude Opus 5 (1M context) <[email protected]>
sohv added a commit to sohv/TransformerLens that referenced this pull request Aug 5, 2026
…input

TransformerBridge.forward() did not derive position_ids from a supplied
attention_mask, so masked-out tokens silently shifted the absolute position of
every real token after them — no error, no NaN, just wrong logits and a wrong
loss. On gpt2 the loss for one prompt moved from 4.503170 unpadded to 11.154946
with three left pads, while HookedTransformer stays invariant (drift ~1e-06).
Right padding was never affected, since causality already protects it.

transformer_bridge.py derived position_ids only for batched *list* input, so
pre-tokenized tensors fell through to HF's plain arange and the offset was never
removed. This reuses utils.get_offset_position_ids — the same helper PosEmbed
and AbstractAttention already use — so the bridge shares HookedTransformer's
position derivation rather than paralleling it. An explicitly supplied
position_ids still wins, and an all-ones mask reduces to arange, so this is a
no-op when there is no padding.

The derivation is offset by any cached prefix: with past_key_values the mask
spans past+new while input_ids holds only the new tokens, so positions are
sliced back to the tokens actually being passed.

The bridge was also inconsistent with itself before this — the same batch gave
different logits depending on whether it was passed as strings or token IDs
(max |logit diff| 4.142e+01) — and enable_compatibility_mode(), which documents
"HookedTransformer-equivalent numerics", diverged on left-padded input while
matching exactly on unpadded input.

Adds integration regression tests: logit invariance under both padding sides,
the same property in compatibility mode, agreement with the shared helper's
derivation, precedence of an explicit position_ids, interior mask gaps, and a
cached decode step. Right-padding cases are controls that pass with and without
the fix. They live in the integration tier because left padding produces a fully
masked query row, which the Native attention path turns into NaN until the
masked-softmax fix in TransformerLensOrg#1608 lands.

Fixes TransformerLensOrg#1609.

Co-Authored-By: Claude Opus 5 (1M context) <[email protected]>
sohv added a commit to sohv/TransformerLens that referenced this pull request Aug 5, 2026
…input

TransformerBridge.forward() did not derive position_ids from a supplied
attention_mask, so masked-out tokens silently shifted the absolute position of
every real token after them — no error, no NaN, just wrong logits and a wrong
loss. On gpt2 the loss for one prompt moved from 4.503170 unpadded to 11.154946
with three left pads, while HookedTransformer stays invariant (drift ~1e-06).

transformer_bridge.py derived position_ids only for batched *list* input, so
pre-tokenized tensors fell through to HF's plain arange and the offset was never
removed. This reuses utils.get_offset_position_ids — the same helper PosEmbed
and AbstractAttention already use — so the bridge shares HookedTransformer's
position derivation rather than paralleling it. An explicitly supplied
position_ids still wins.

The derivation fires only when the mask actually moves an attended token off its
default position, i.e. when some masked token precedes a real one. That covers
left padding and interior mask gaps. Pure right padding and all-ones masks
already agree with arange, so they are left alone: injecting position_ids there
is a no-op at best, and breaks models whose forward does not accept the argument
or which compute their own position streams (multimodal mRoPE).

With a KV cache the mask spans past+new while input_ids holds only the new
tokens, so the derived positions are sliced back to the tokens being passed.

The bridge was also inconsistent with itself before this — the same batch gave
different logits depending on whether it was passed as strings or token IDs
(max |logit diff| 4.142e+01) — and enable_compatibility_mode(), which documents
"HookedTransformer-equivalent numerics", diverged on left-padded input while
matching exactly on unpadded input.

Adds integration regression tests: logit invariance under both padding sides,
the same property in compatibility mode, agreement with the shared helper,
precedence of an explicit position_ids, interior mask gaps, a cached decode
step, and that no position_ids are injected when the mask does not require it.
Right-padding cases are controls that pass with and without the fix. They live
in the integration tier because left padding produces a fully masked query row,
which the Native attention path turns into NaN until the masked-softmax fix in
TransformerLensOrg#1608 lands.

Fixes TransformerLensOrg#1609.

Co-Authored-By: Claude Opus 5 (1M context) <[email protected]>

@jlarson4 jlarson4 left a comment

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.

Thanks for reporting this and putting together a fix! There are a couple tricky areas with this mask, I put a couple notes in. Let me know if you have any questions.

Additionally, in remote_bridge.py, there is another calcite of this helper and it still passes no attention_mask, so RemoteBridge.forward(..., return_type="loss") still has the bug. The mask is already in scope in **kwargs. Could you add attention_mask=kwargs.get("attention_mask") to the call at line 113?

# (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)

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.

Swapping in-place *= for masked_fill drops an implicit length guard: *= raised when the mask was longer than the scored tokens, but masked_fill broadcasts both operands and replicates the value instead. Now that the mask is forwarded, bridge(new_tokens, attention_mask=<cache+new>, past_key_values=cache, return_type="loss", loss_per_token=True) returns shape (1,5) with one value repeated five times, but previously it returned (1,1), and the scalar variant raises. Could you add an explicit attention_mask.shape[1] == tokens.shape[1] assertion here so the mismatch fails loudly again?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Addressed in df5f46d. lm_cross_entropy_loss now explicitly asserts that attention_mask.shape == tokens.shape, while BridgeCore first slices or reduces valid cache plus new masks to the scored token window. I also added mismatch and cached-window regression coverage.

Comment thread transformer_lens/utilities/lm_utils.py Outdated
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"]]:

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.

pos_minus_one is a fresh jaxtyping name bound only by the return value, so it matches any length including pos, the exact off-by-one it exists to catch. The repo already spells this axis symbolically as "batch pos-1", which binds against pos from the arguments and so actually catches it. Can you use that form to match?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Addressed in df5f46d. The per-token loss return annotation now uses batch pos-1, binding the result length to the input pos dimension.

Comment thread transformer_lens/utilities/lm_utils.py Outdated

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.

lm_accuracy's per_token=True branch is annotated Float[torch.Tensor, "batch pos"] but returns a bool tensor of shape [batch, pos-1]. It is wrong in both dtype and shape, and its own docstring says [batch, seq_len-1]. It raises BeartypeCallHintReturnViolation today under the repo's pytest config, and it is public API via transformer_lens.utils, reachable from Bridge caches. Since you're already correcting the sibling function in this file, could you fix it here too?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Addressed in df5f46d. lm_accuracy(per_token=True) is now annotated as Bool[Tensor, batch pos-1], and a runtime test checks both its boolean dtype and [batch, pos-1] shape.

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(

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 Bridge's native forward documents 2-D and 4-D masks as equivalent (sources/native/model.py:344-349), and they are. On the issue's own tiny bridge, a 4-D bool encoding of the same padding gives bit-identical logits to the 2-D form. But the loss path now distinguishes them: at this head the 2-D call returns the correct 3.8386409282684326 while the equivalent 4-D call raises RuntimeError: The size of tensor a (6) must match the size of tensor b (5) (on dev-4.x it returned a number). The same unvalidated forwarding is what turns the KV-cache case in lm_utils.py on line 40 from a number into a crash. Could you slice/reduce the mask to the scored [batch, pos] window before loss_fn?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Addressed in df5f46d. BridgeCore now normalizes 2-D, 4-D bool or additive, and cache plus new masks to the scored [batch, pos] window before calling loss_fn. Tests cover key-only and full-causal 4-D masks, bit-identical 2-D or 4-D logits, equal manually anchored losses, and rectangular cached 4-D masks.

assert torch.count_nonzero(loss[~next_token_mask]) == 0


def test_forward_loss_is_finite_with_left_padding() -> None:

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.

Every assertion in the file is bridge-vs-bridge. A masked loss with the wrong denominator or mask shift would satisfy all four tests, this function in particular is value-blind. And the lm_utils.py:40 NaN-safety hunk has no test. As a regression check I reverted it to *= and all 4 tests stayed green (the native hunk alone keeps the logits finite). Could you add one anchor against an externally-computed value, staying inside the Bridge?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Addressed in df5f46d. The Bridge regression now computes the expected value independently with torch.nn.functional.cross_entropy over manually selected valid transitions instead of calling bridge.loss_fn. I also added a direct NaN-masking regression that fails with multiplication and passes with masked_fill.

@emerardd

emerardd commented Aug 6, 2026

Copy link
Copy Markdown
Contributor Author

Thanks for the detailed review. I pushed df5f46d addressing all five inline threads. The review-body RemoteBridge issue is also fixed: RemoteBridge.forward now passes attention_mask=kwargs.get("attention_mask") into _finalize_return, with a cache-plus-new masked-loss regression. Validation: 40 focused tests passed, formatting checks passed, and uv run mypy . reports no issues in 422 source files. I left the threads unresolved for reviewer verification.

@jlarson4

jlarson4 commented Aug 6, 2026

Copy link
Copy Markdown
Collaborator

Awesome, thank you @emerardd. I will do a re-review shortly.

@jlarson4

jlarson4 commented Aug 6, 2026

Copy link
Copy Markdown
Collaborator

Great work on this @emerardd, no additional change requests on my end, merging now!

@jlarson4
jlarson4 merged commit 7ebeab9 into TransformerLensOrg:dev-4.x Aug 6, 2026
25 checks passed
jlarson4 pushed a commit that referenced this pull request Aug 6, 2026
…input (#1610)

* fix(bridge): derive position_ids from attention_mask for left-padded input

TransformerBridge.forward() did not derive position_ids from a supplied
attention_mask, so masked-out tokens silently shifted the absolute position of
every real token after them — no error, no NaN, just wrong logits and a wrong
loss. On gpt2 the loss for one prompt moved from 4.503170 unpadded to 11.154946
with three left pads, while HookedTransformer stays invariant (drift ~1e-06).

transformer_bridge.py derived position_ids only for batched *list* input, so
pre-tokenized tensors fell through to HF's plain arange and the offset was never
removed. This reuses utils.get_offset_position_ids — the same helper PosEmbed
and AbstractAttention already use — so the bridge shares HookedTransformer's
position derivation rather than paralleling it. An explicitly supplied
position_ids still wins.

The derivation fires only when the mask actually moves an attended token off its
default position, i.e. when some masked token precedes a real one. That covers
left padding and interior mask gaps. Pure right padding and all-ones masks
already agree with arange, so they are left alone: injecting position_ids there
is a no-op at best, and breaks models whose forward does not accept the argument
or which compute their own position streams (multimodal mRoPE).

With a KV cache the mask spans past+new while input_ids holds only the new
tokens, so the derived positions are sliced back to the tokens being passed.

The bridge was also inconsistent with itself before this — the same batch gave
different logits depending on whether it was passed as strings or token IDs
(max |logit diff| 4.142e+01) — and enable_compatibility_mode(), which documents
"HookedTransformer-equivalent numerics", diverged on left-padded input while
matching exactly on unpadded input.

Adds integration regression tests: logit invariance under both padding sides,
the same property in compatibility mode, agreement with the shared helper,
precedence of an explicit position_ids, interior mask gaps, a cached decode
step, and that no position_ids are injected when the mask does not require it.
Right-padding cases are controls that pass with and without the fix. They live
in the integration tier because left padding produces a fully masked query row,
which the Native attention path turns into NaN until the masked-softmax fix in
#1608 lands.

Fixes #1609.

Co-Authored-By: Claude Opus 5 (1M context) <[email protected]>

* fix(bridge): gate derived position_ids on the target model, decide per row

The mask-derived position_ids from the previous commit were handed to every
model that hit the predicate. Two families break on that:

  * Fixed-signature remote-code forwards (LLaDAModelLM) take neither
    position_ids nor **kwargs, so a left-padded input raised TypeError where
    the model used to return logits.
  * Models that own their position derivation. mRoPE architectures build a
    3-D index in get_rope_index, but only while position_ids is None; a
    supplied 2-D tensor is silently broadcast across all three streams.
    Left-padded multimodal input drifted 1.414e-04 off fresh HF.
  * OPT's positional embedding consumes the mask and derives the same
    positions, using its own convention for the padded slots, so injecting
    replaced a correct derivation with a differing one.

_accepts_derived_position_ids() refuses all three, mirroring how
output_attentions is guarded in BridgeCore.run_with_cache. Their own
derivations already place positions on attended slots only, so deferring is
correct rather than a gap.

The predicate is now per row instead of a whole-batch .any(): rows that are
unpadded or purely right-padded keep arange verbatim, so one left-padded row
in a batch no longer perturbs its neighbours (8.223e-01 -> 0 on distilgpt2).

Also: cast the mask to long before deriving, since a float 0/1 mask produced
float positions and crashed the embedding lookup.

Cached decoding under left padding is fixed as a consequence -- prefilling
through the bridge and stepping now lands 1.755e-04 from the unpadded result,
against 2.844e+01 before. The previous test prefilled with raw HF and asserted
only shape and finiteness, so it modelled the one pattern that stays wrong.

Co-Authored-By: Claude Opus 5 <[email protected]>

---------

Co-authored-by: Claude Opus 5 (1M context) <[email protected]>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants