Repository navigation
Fix non-finite gradients for left-padded batches in compatibility mode - #1811
Conversation
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
left a comment
There was a problem hiding this comment.
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.
|
Good solution to this issue @aliezaway, looks great, merging now |
Description
Fixes non-finite gradients from a compatibility-mode
TransformerBridgeon 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-infand softmax returns NaN._scrub_compatibility_pattern_nanszeroes 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 withtorch.where, which zeroed those gradients, so it never showed this.Fix. A new
AttentionBridge._masked_softmaxhelper. 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_scoresstill shows-inffor masked positions. The existing NaN scrub is kept for HookedTransformer parity.Fixes #1809
Type of change
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=Truecompatibility 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
devand 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.
devJointQKVAttentionBridgeJointQKVPositionEmbeddingsAttentionBridgeBloomAttentionBridgePositionEmbeddingsAttentionBridgePositionEmbeddingsAttentionBridge* 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
devand this branch.A tiny random Gemma2 (
PositionEmbeddingsAttentionBridgewith softcap) gives the same result: NaN ondev, finite with this PR, drift 1.8e-6.Local suites (macOS arm64, Python 3.12, torch 2.11):
make check-formatanduv 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_searchfails identically ondevon this machine: LBFGS takes thetolerance_gradexit instead ofline_search.test_get_params[google/gemma-2-2b-it]andtest_TransformerBridge_gemma2_forwardfail with a 401 while fetching the gated repo'sconfig.json: I have noHF_TOKENlocally, so no model code ran.test_mistral3_adapter.pytest hung on a stalled HF download; the file passes on rerun (6 passed).Checklist: