Skip to content

Compatibility-mode bridge returns non-finite gradients on a left-padded, masked batch (4.0.0) #1809

Description

@speediedan

Human Preface (@speediedan)

Thanks again for TransformerLens! Another follow-on to the issues from earlier this week as I bump support in
(interpretune.org).

Basically, our gradient conformance cases went non-finite on the compatibility-mode bridge.

I've reviewed and human-validated the rest of this description below but it is largely
Opus (5.5) written. Let me know if you want more human clarification/feedback/revision of it.
Thanks!

Summary

In 4.0.0, a TransformerBridge in compatibility mode, given a left-padded batch together with its
attention_mask, returns finite logits but non-finite gradients. The non-finite values are not confined to
the pad positions: every real token of the padded row gets them too, at every layer. The same call without
compatibility mode is finite, the same call with right padding is finite, and 3.5.1 (the v3.5.1 tag) is finite
for the identical input.

Reproduction

import torch
import transformer_lens.model_bridge.sources.transformers  # attaches boot_transformers
from transformer_lens.model_bridge import TransformerBridge
from transformers import AutoTokenizer

tok = AutoTokenizer.from_pretrained("gpt2")
tok.pad_token = tok.eos_token
tok.padding_side = "left"
batch = tok(
    ["The capital of France is", "A much longer prompt about the capital city of a European country is"],
    return_tensors="pt",
    padding=True,
)

bridge = TransformerBridge.boot_transformers("gpt2", device="cpu")
bridge.enable_compatibility_mode()

kept = {}

def keep(tensor, hook):
    tensor.retain_grad()
    kept["resid"] = tensor
    return tensor

logits = bridge.run_with_hooks(
    batch.input_ids, attention_mask=batch.attention_mask, fwd_hooks=[("blocks.0.hook_in", keep)]
)
logits.sum().backward()
print(torch.isfinite(logits).all())              # tensor(True)
print(torch.isfinite(kept["resid"].grad).all())  # tensor(False)

What was measured

gpt2, CPU, transformer-lens 4.0.0, torch 2.13.0, transformers 5.17.0, Python 3.13. The probe imports
only transformer_lens, transformers and torch.

padding configuration forward finite gradient finite
left no compatibility mode yes yes
left enable_compatibility_mode() yes no
left enable_compatibility_mode(no_processing=True) yes no
left fold_ln only yes no
left center_writing_weights only yes no
right each of the five above yes yes

Because no_processing=True fails too, the trigger looks like the compatibility-mode path itself, not any
weight-processing step. At blocks.0.hook_in the gradient is non-finite at every position of the padded row (all 8
pads and all 5 real tokens), while the unpadded row in the same batch stays finite; every other layer's hook_in
gradient is non-finite as well. On 3.5.1 the same script prints True for both.

A guess at where to look

Not verified. With left padding, the query at a row's first position can attend only to pad keys, so its
attention row is fully masked. A softmax over a fully masked row is finite only if the masking avoids -inf
everywhere or the row is handled explicitly. The forward stays finite, so if this is the cause, something
downstream is replacing those values and the backward is not.

Why it matters

Gradient-based attribution over batched prompts commonly uses left padding (the side generation uses), and here
the corruption reaches the real tokens, so every attribution score for a padded row is invalid. Nothing
raises, so the error surfaces only as non-finite or implausible scores.

Prior art

Searched the issues for NaN gradients and for left padding. #1609 (closed) made the bridge derive
position_ids from the attention mask; it concerns forward values, not gradients. No open or closed issue
covers this.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    TransformerBridgeBug specific to the new TransformerBridge systembugSomething isn't workingcomplexity-simpleSimple issues, which may be good for beginners

    Type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions