On M-series decode, sdpa_vector_2pass_1 assigns one simdgroup per query head, and the gqa_factor simdgroups of a threadgroup all stream the same K/V block (the keys/values offsets don't depend on the simdgroup index). Loads therefore scale with query heads, not kv heads, and for models with few kv heads the kernel saturates load-issue before DRAM.
Measured effective bandwidth of mx.fast.scaled_dot_product_attention (M5 Pro 24GB, fp16, q_len 1, KV bytes counted once, distinct KV buffers cycled between calls to keep SLC out of the numbers):
| geometry |
ctx 2K |
ctx 8K |
ctx 32K |
| 32q/32kv (gqa 1) |
239 GB/s |
282 GB/s |
281 GB/s |
| 32q/8kv (gqa 4) |
234 GB/s |
263 GB/s |
264 GB/s |
| 32q/4kv/hd128 (gqa 8, Qwen3-30B/235B) |
130 GB/s |
184 GB/s |
190 GB/s |
| 64q/8kv/hd64 (gqa 8, gpt-oss-20b) |
76 GB/s |
113 GB/s |
119 GB/s |
A dense fp16 qmv on the same machine streams at 266 GB/s, so the gqa-8 rows are kernel-shape, not hardware.
At 32K context this is most of the decode step: Qwen3-30B-A3B-4bit (truncated to 24 layers so 32K KV fits in 24 GB) spends ~60% of its 15 ms/step streaming KV at this reduced rate.
Reassigning the work so that each simdgroup owns a contiguous token sub-chunk and computes several of its group's query heads from registers (so each K/V byte is read gqa_factor / heads_per_simd times instead of gqa_factor) recovers most of the gap: 1.14-1.27x kernel-level at 32K, and +9.5% / +9.2% end-to-end on Qwen3-30B-A3B and Qwen2.5-3B, with the second pass unchanged.
On M-series decode,
sdpa_vector_2pass_1assigns one simdgroup per query head, and thegqa_factorsimdgroups of a threadgroup all stream the same K/V block (thekeys/valuesoffsets don't depend on the simdgroup index). Loads therefore scale with query heads, not kv heads, and for models with few kv heads the kernel saturates load-issue before DRAM.Measured effective bandwidth of
mx.fast.scaled_dot_product_attention(M5 Pro 24GB, fp16, q_len 1, KV bytes counted once, distinct KV buffers cycled between calls to keep SLC out of the numbers):A dense fp16 qmv on the same machine streams at 266 GB/s, so the gqa-8 rows are kernel-shape, not hardware.
At 32K context this is most of the decode step: Qwen3-30B-A3B-4bit (truncated to 24 layers so 32K KV fits in 24 GB) spends ~60% of its 15 ms/step streaming KV at this reduced rate.
Reassigning the work so that each simdgroup owns a contiguous token sub-chunk and computes several of its group's query heads from registers (so each K/V byte is read
gqa_factor / heads_per_simdtimes instead ofgqa_factor) recovers most of the gap: 1.14-1.27x kernel-level at 32K, and +9.5% / +9.2% end-to-end on Qwen3-30B-A3B and Qwen2.5-3B, with the second pass unchanged.