Skip to content

sdpa_vector_2pass_1 loads K/V once per query head instead of once per kv head #4076

Description

@dudududukim

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.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions