Skip to content

perf: reduce temporary memory when computing attention head results - #1854

Merged
jlarson4 merged 3 commits into
TransformerLensOrg:devfrom
kiteretsu903:perf-bound-head-result-memory
Oct 7, 2026
Merged

jlarson4 merged 3 commits into
TransformerLensOrg:devfrom
kiteretsu903:perf-bound-head-result-memory

Conversation

@kiteretsu903

Copy link
Copy Markdown

Description

I noticed that ActivationCache.compute_head_results() creates a large temporary tensor before summing over each head's hidden dimension. For one GPT-2-sized layer with 512 tokens, that temporary takes about 1,208 MB, even though the result is only about 19 MB. This can make head-level analysis unnecessarily expensive on longer inputs.

I changed this to process the input in chunks. When weight gradients aren't needed, I split over tokens. When they are needed, I split over attention heads so each weight gradient still sums over all tokens at once. This avoids overflowing FP16 partial gradients before positive and negative contributions can cancel. The calculation still uses the original multiply-and-sum operations, with no API or dependency changes.

In a local float32 CPU benchmark with 512 tokens, 12 heads, d_head=64 and d_model=768:

Mode Largest temporary before Largest temporary after Median forward time before Median forward time after
No weight gradients 1,208 MB 33 MB 314 ms 66 ms
Weight gradients enabled 1,208 MB 101 MB 200 ms 247 ms

These are three-trial measurements with four CPU threads. The weight-gradient path saves memory but was slower in this benchmark. Each chunk retains at least one whole token or head, so its size can exceed the target for larger inputs.

Type of change

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

Validation

I added three test functions covering 13 cases: outputs and gradients across four floating-point types, non-contiguous inputs, temporary allocations, and FP16 gradient cancellation. All 13 pass. On the original implementation, the nine numerical cases pass and the four allocation checks fail.

Using Python 3.12.14 and PyTorch 2.11.0 CPU in the existing environment with UV_NO_SYNC=1:

  • ActivationCache tests: 153 passed.
  • make format and uv run mypy .: passed, with no typing issues in 387 source files.
  • make test-pr, with CI=true and four pytest workers: 8,636 passed, 108 skipped, 3 expected failures. The expected failures are existing hook-alias cases; I haven't added any skip or xfail markers.
  • A real GPT-2 run with 512 tokens: all 12 layers and the stacked head results exactly matched the original implementation.

Checklist:

  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation: updated the method docstring.
  • My changes generate no new warnings: the test suite still reports dependency and model 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

@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 the careful measurements on this performance improvement @kiteretsu903! Just one note on the mechanism & one on the tests, let me know if you have any questions

Comment thread transformer_lens/ActivationCache.py Outdated
Comment thread tests/unit/test_activation_cache.py Outdated
@kiteretsu903

Copy link
Copy Markdown
Author

@jlarson4

thanks for the suggestions! I tried the bridge's einsum approach and it works much better than chunking. I've switched over to that and kept dtype promotion for mixed-dtype caches.

I updated the tests to compare outputs and gradients against a separate float64 reference, with tolerances for each dtype. I also tightened the allocation checks for contiguous, same-dtype inputs. The FP16 cancellation test still checks for an exact zero.

I did hit one issue on an L4: at Llama widths, two FP32 activation-gradient values fell outside the existing tolerance. Per-head matmuls passed with the same inputs and upstream gradient, so I've used those for CUDA FP32 activation gradients when autocast is off. Everything else keeps einsum, and the tolerances are unchanged.

here's what passed:

Check Result
ActivationCache CPU tests 160 passed
Related acceptance/integration tests 51 passed
CUDA numerical checks 48 passed
Extra CUDA cases: seeds, empty shapes, broadcasting and second derivatives 10 passed
Formatting and full-project mypy Passed

The GPU checks cover FP32, FP16, BF16, autocast, mixed dtypes, batched/batchless inputs and gradients.

For synthetic Llama-width weights, 32 tokens and both gradients on the L4:

Implementation Median time Additional peak tensor memory
Original PR chunking 47.10 ms 216.53 MB
Revised code 3.99 ms 151.52 MB

That's for a single layer's forward and backward. The FP32 fallback takes more time and memory than plain einsum, but still improves on the original chunking and passes the accuracy check.

I also finished the non-slow CPU integration run: 1,426 passed, six skipped and one optimizer-parity failure. I reran that one with the original and revised methods and got the same result both times, with zero calls to compute_head_results. I suggest opening a separate issue to investigate this

@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 the quick adjustments @kiteretsu903! Good coverage on all my comments. A couple additional comments related to the CUDA system you added

Comment thread tests/unit/test_activation_cache.py Outdated
actual_grads = torch.autograd.grad(actual, (z, weights), grad_output)
torch.testing.assert_close(actual.double().cpu(), expected, rtol=2e-5, atol=2e-4)
for actual_grad, expected_grad in zip(actual_grads, expected_grads):
torch.testing.assert_close(actual_grad.double().cpu(), expected_grad, rtol=2e-5, atol=2e-4)

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.

This test fails on macOS arm64: the z-gradient is an fp32 reduction over d_model=4096 with entries up to a few hundred, and rtol 2e-5 / atol 2e-4 is tighter than fp32 matmul reductions guarantee, so the result depends on the BLAS backend. The unit tier runs on macos-latest in the MPS workflow for merges to main, where this would be red. Can the gradient tolerance be widened to what fp32 reductions actually guarantee across backends?

Comment thread transformer_lens/ActivationCache.py Outdated
and not torch.is_autocast_enabled("cuda")
and z.shape[-2] == weights.shape[0] != 0
):
# Per-head GEMMs avoid less accurate CUDA batched activation-gradient reductions.

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.

No CI runner has CUDA, so this branch never executes in the test suite, and it makes cached head results bitwise-diverge from the forward pass's batched contraction on CUDA fp32 training setups. On CPU the per-head and batched forms produce identical gradients, which suggests the difference being compensated for is backend rounding of the same kind the test above is measuring. We may want to drop this branch unless there's a reproducible CUDA case that misses a realistic tolerance. Would you be willing to investigate and resolve that?

@kiteretsu903

Copy link
Copy Markdown
Author

@jlarson4

thanks for catching this! I reproduced the failure on an M4 with macOS Accelerate and tested a local revision that removes the CUDA FP32 per-head fallback and uses einsum throughout.

I widened only the FP32 activation-gradient tolerance to rtol=1e-4 / atol=1e-3. Output and weight-gradient tolerances are unchanged. I also added a normalized error check against the float64 reference, scaled by the sum of absolute products, to cover cancellation and near-zero results. Additional seeds and the original L4 case passed without the fallback. I didn’t find a reproducible CUDA case in these checks that required the fallback under the revised tolerance, so I removed it.

here's what passed:

  • M4 ActivationCache CPU tests: 161 passed
  • Related acceptance/integration tests: 51 passed
  • L4 CUDA numerical checks: 48 passed
  • Extra CUDA cases: 10 passed
  • Formatting and full-project mypy: passed

For synthetic Llama-width weights, 32 tokens and both FP32 gradients on the L4, plain einsum took 1.37 ms and 84.41 MB of additional peak tensor memory, compared with 9.87 ms and 151.52 MB for the fallback in the same run. These are single-layer forward/backward measurements.

The non-slow CPU integration run again had 1,426 passes and the same optimizer-parity failure. The baseline and local revision produced exactly the same failure, with zero calls to compute_head_results.

@jlarson4

jlarson4 commented Oct 7, 2026

Copy link
Copy Markdown
Collaborator

Looks great @kiteretsu903, thank you for making that adjustment so swiftly! Approved and merging now

@jlarson4
jlarson4 merged commit 0033815 into TransformerLensOrg:dev Oct 7, 2026
27 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.

2 participants