Skip to content

Adding VJP for the scaled dot product attention. - #4563

Open
tpegolotti wants to merge 11 commits into
mainfrom
sdpa-vjp
Open

tpegolotti wants to merge 11 commits into
mainfrom
sdpa-vjp

Conversation

@tpegolotti

@tpegolotti tpegolotti commented Sep 25, 2026 •

Copy link
Copy Markdown
Collaborator

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.

sdpa_vjp sdpa_vjp_4

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

@tpegolotti
tpegolotti requested a review from jagrit06 September 25, 2026 11:46
@tpegolotti
tpegolotti marked this pull request as ready for review September 28, 2026 07:46
@tpegolotti
tpegolotti marked this pull request as draft September 28, 2026 07:48
@tpegolotti
tpegolotti marked this pull request as ready for review September 28, 2026 09:26
@tpegolotti
tpegolotti requested a review from zcbenz September 28, 2026 09:26

@zcbenz zcbenz left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

It is awesome that this not only massively reduces memory use, but also has a big speedup!

Comment thread mlx/backend/metal/scaled_dot_product_attention.cpp Outdated
@tpegolotti
tpegolotti requested a review from zcbenz September 28, 2026 17:40

This branch has not been deployed

No deployments
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