Repository navigation
perf: reduce temporary memory when computing attention head results - #1854
Conversation
jlarson4
left a comment
There was a problem hiding this comment.
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
|
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:
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:
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
left a comment
There was a problem hiding this comment.
Thanks for the quick adjustments @kiteretsu903! Good coverage on all my comments. A couple additional comments related to the CUDA system you added
| 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) |
There was a problem hiding this comment.
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?
| 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. |
There was a problem hiding this comment.
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?
|
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:
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. |
|
Looks great @kiteretsu903, thank you for making that adjustment so swiftly! Approved and merging now |
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=64andd_model=768: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
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:make formatanduv run mypy .: passed, with no typing issues in 387 source files.make test-pr, withCI=trueand 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.Checklist: