Adding VJP for the scaled dot product attention. - #4563
Open
tpegolotti wants to merge 11 commits into
Open
tpegolotti wants to merge 11 commits into
tpegolotti wants to merge 11 commits into
Conversation
tpegolotti
marked this pull request as ready for review
September 28, 2026 07:46
tpegolotti
marked this pull request as draft
September 28, 2026 07:48
tpegolotti
marked this pull request as ready for review
September 28, 2026 09:26
zcbenz
approved these changes
Sep 28, 2026
zcbenz
left a comment
Member
There was a problem hiding this comment.
It is awesome that this not only massively reduces memory use, but also has a big speedup!
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Adding an implementation for the scaled dot product attention VJP, needed for fine-tuning and training.
Implementation
On Metal the unified memory architecture makes bandwidth less of a bottleneck than it is on standard GPUs, so the VJP fallback is already fast. However, it uses a lot of memory since it allocates a full [context_size x context_size] attention matrix.
Writing a specialized kernel did not improve performance since it requires more compute: the attention matrix is never materialized and instead is computed, implicitly, twice. Expressing the VJP instead as a sequence of tiled matmuls, using steel, gives a large memory reduction and speedup.
Results
benchmarks/python/sdpa_bench.py now takes a -bw flag to benchmark the VJP against the fallback on both time and memory. Plots below are M5 Ultra, 80 cores, batch 1 and 4, sweeping context size with one line per (head_dim, n_qh, n_kvh) and line style separating causal from unmasked.
Here are the results for batch=1 and batch=4. The trends are similar in both plots. For the causal case we have up to 2.8x speedups. Memory improves in every configuration, up to 24x at the longest shape: 24.63 GiB to 1 GiB for batch size = 1, and from 97GiB to 4GiB for batch size = 4.
The exception is (256, 24, 4), the grey line, which can be slower than the fallback. Memory there is still better, so the tradeoff looks worth it.
Edit: I've removed non-needed buffers and improved memory and speedup further: 3a2abe7 and be9ef2f