Skip to content

Fix non-finite gradients for left-padded batches in compatibility mode - #1811

Merged
jlarson4 merged 1 commit into
TransformerLensOrg:devfrom
aliezaway:fix/compat-left-pad-nan-grad
Sep 25, 2026
Merged

jlarson4 merged 1 commit into
TransformerLensOrg:devfrom
aliezaway:fix/compat-left-pad-nan-grad

Conversation

@aliezaway

Copy link
Copy Markdown
Contributor

Description

Fixes non-finite gradients from a compatibility-mode TransformerBridge on left-padded, masked batches.

Root cause. Compatibility mode masks attention scores with -inf (#1694). With left padding, the leading pad queries of a padded row see no visible key, so their whole score row is -inf and softmax returns NaN. _scrub_compatibility_pattern_nans zeroes that NaN, but only in the forward pass. Softmax's backward still multiplies by the saved NaN output (y * (g - Σ g·y)), and the key-padding mask is applied by addition, which passes gradients straight through. So the NaN reached Q/K, the residual stream and then every position of the padded row at every layer. HookedTransformer 3.5.1 applied all masking with torch.where, which zeroed those gradients, so it never showed this.

Fix. A new AttentionBridge._masked_softmax helper. In compatibility mode it replaces fully masked rows with zeros before softmax and zeroes their pattern afterwards. Those rows are therefore finite in both passes and still produce the all-zero pattern that HookedTransformer produced. Outside compatibility mode it is plain softmax, bit for bit. It is used in _softmax_dropout_pattern, which covers most attention bridges, and at the two sites that call softmax directly: PositionEmbeddingsAttentionBridge (non-sink path) and the LLaDA attention bridge. hook_attn_scores still shows -inf for masked positions. The existing NaN scrub is kept for HookedTransformer parity.

Fixes #1809

Type of change

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

Verification

The issue's reproduction (gpt2, left padding, enable_compatibility_mode()) now gives finite gradients for both rows. Right padding and non-compat mode are unchanged.

The new integration test checks more than finiteness. Real tokens never attend to pads, so the gradients at the padded row's real positions must match the same prompt run unpadded, and the pad positions must receive exactly zero gradient. It runs with both processed and no_processing=True compatibility mode.

The test uses a logsumexp loss instead of logits.sum(): with processed weights the unembed is centered, so summed logits hardly depend on the residual stream and would make the check nearly vacuous. The remaining padded-vs-unpadded difference is fp32 accumulation noise (relative norm ≈ 2e-6). The non-compat HF path, which this PR does not touch, drifts by exactly the same amount.

Both new tests fail on dev and pass with the fix.

Beyond GPT-2, I ran the same left-padding gradient check (compatibility mode, gradient of the first block's input) on one model per attention bridge. The check is: finite gradients, zero gradient at pad positions, and padded real-token gradients equal to the unpadded ones.

Model Attention bridge dev This PR (fp32 drift) This PR (fp64 drift)
gpt2 JointQKVAttentionBridge NaN 1.4e-6 –
pythia-70m JointQKVPositionEmbeddingsAttentionBridge NaN 5.1e-6 5.6e-15
bigscience-small-testing BloomAttentionBridge NaN 4.0e-4 1.7e-14
tiny-random-Llama PositionEmbeddingsAttentionBridge NaN 9.8e-6 0
tiny-Qwen2-2.5 PositionEmbeddingsAttentionBridge NaN 1.0e-3 * –

* The tiny random Qwen2 has near-uniform logits. With the unembed centered, the logsumexp gradient is about 500× smaller than without centering, so relative rounding error is amplified. With a single-token-logit loss the drift is 9.8e-8, the same as the untouched HF path (1.3e-7). The forward logits are padding-invariant to 1e-7 in every mode, identically on dev and this branch.

A tiny random Gemma2 (PositionEmbeddingsAttentionBridge with softcap) gives the same result: NaN on dev, finite with this PR, drift 1.8e-6.

Local suites (macOS arm64, Python 3.12, torch 2.11):

  • unit: 6613 passed
  • integration: 1370 passed
  • acceptance: 162 passed
  • docstring: 13 passed
  • make check-format and uv run mypy . are clean.

Tests that did not pass locally, none of them related to this change:

  • tests/unit/tools/test_sparse_probing.py::test_stop_reason_reports_a_stalled_line_search fails identically on dev on this machine: LBFGS takes the tolerance_grad exit instead of line_search.
  • test_get_params[google/gemma-2-2b-it] and test_TransformerBridge_gemma2_forward fail with a 401 while fetching the gated repo's config.json: I have no HF_TOKEN locally, so no model code ran.
  • One test_mistral3_adapter.py test hung on a stalled HF download; the file passes on rerun (6 passed).

Checklist:

  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation (none needed)
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

Compatibility mode masks attention scores with -inf, so the leading pad
queries of a left-padded row have no visible key and softmax to NaN.
The NaN was scrubbed from the forward pattern, but softmax's backward
still multiplied by the saved NaN output, and the additive padding mask
passed it through to Q/K, so every position of the padded row got
non-finite gradients.

Add AttentionBridge._masked_softmax: in compatibility mode it zeroes
fully masked rows before softmax and zeroes their pattern afterwards,
keeping both passes finite while still producing HookedTransformer's
all-zero pattern. Outside compatibility mode it is plain softmax. Use
it in _softmax_dropout_pattern and at the two direct softmax sites
(PositionEmbeddingsAttentionBridge and the LLaDA attention bridge).

Fixes TransformerLensOrg#1809

@koriyoshi2041 koriyoshi2041 left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Validated at 733a932b. The compatibility-only branch removes fully masked rows from softmax's saved output, then restores the zero attention pattern, so both the forward contract and backward gradient are well defined. The shared helper is also wired through the general, position-embedding, and LLaDA softmax paths while non-compatibility mode remains a direct softmax.

The focused synthetic backward tests pass locally (2/2), including zero gradients for the fully masked row and bit-for-bit plain-softmax behavior outside compatibility mode. The exact-head hosted coverage, type, format, benchmark, notebook, and Python 3.10–3.12 compatibility checks are green. I also started the GPT-2 integration file locally, but model download did not complete in five minutes, so I did not count that as a local pass.

@jlarson4

Copy link
Copy Markdown
Collaborator

Good solution to this issue @aliezaway, looks great, merging now

@jlarson4
jlarson4 merged commit 5ca0bea into TransformerLensOrg:dev Sep 25, 2026
26 checks passed
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.

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

3 participants