<!-- llms-explorer concept facts · https://llms-explorer.com/tree/metal-flash-attention-kernel-partitioning-differ/ · pack 2026-10-05 · ~4156 tokens -->

# Metal flash-attention kernel partitioning differences across M-series chips

> Chip class comes from the last character of `d.get_architecture()`; the device constructor comments map 'p' to phone, 'g' to base or Pro, 's' to Max and 'd' to Ultra, and the generation is parsed from the two characters before it.

Parent: [Mac local LLMs: Quantization evaluation](https://llms-explorer.com/tree/mac-local-llms-quantization-evaluation/) · 2 facets · 61 facts · page: https://llms-explorer.com/tree/metal-flash-attention-kernel-partitioning-differ/

## Facts

- Chip class comes from the last character of `d.get_architecture()`; the device constructor comments map 'p' to phone, 'g' to base or Pro, 's' to Max and 'd' to Ultra, and the generation is parsed from the two characters before it. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/device.cpp)
- Decode-shaped attention (query length 8 or less) goes to the vector kernels. The 2-pass vector kernel is chosen when the class is 'd' or 's' and the key length is at least 1024, or when there are fewer KV heads than query heads and the key length is at least 4096; every other case runs the 1-pass kernel with a fixed 1024-thread threadgroup. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- On a 'g' or 'p' class chip with plain multi-head attention (as many KV heads as query heads) the 2-pass kernel is therefore never chosen, and with GQA only from 4096 keys; an M5 MacBook Air was reported to switch to 2-pass at 4096 keys where an M3 Max switched at 1024. — [source](https://github.com/ml-explore/mlx/issues/4022)
- The 2-pass kernel splits the key sequence into `blocks` partial tiles that a second kernel reduces; the grid is (kv heads, batch, blocks) and the threadgroup is (32, gqa_factor, query length), one simdgroup per (query head, query row) pair, where `n_simds = gqa_factor * query length`. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- Block count for class 's': 64, raised when the key length N is above 1024 and `n_simds` is above 4 to 128 (N up to 8192), 256 (up to 32768), 512 (up to 65536) and 1024 beyond. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- Block count for class 'd': 128 by default, 256 when `n_simds` is 2 or less and N is above 8192, and, when `n_simds` is 6 or more, 512 for N from 16384 to 65535 and 1024 from 65536. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- Block count for every other class that reaches 2-pass (only via the GQA 4096-key rule): 64 when `n_simds` is 4 or more, else 32. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- An Ultra never uses fewer than 128 blocks while a Max starts at 64, and the two classes disagree for mid-size `n_simds` (3 to 5) at long context; neither table is documented as measured per chip. — source: `asserted`
- `blocks` is a Metal function constant, so each value compiles a separate pass-1 pipeline, and the hash name of the pipeline includes the block count. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- `MLX_SDPA_BLOCKS` overrides the heuristic; positive values are rounded up to a multiple of 32 because pass 2 reduces the partials in 32-wide chunks and silently drops a tail otherwise. — [source](https://ml-explore.github.io/mlx/build/html/usage/environment_variables.html)
- A specialized pass-1 kernel `sdpa_vector_2pass_1_gqa` replaces the generic one only for single-token queries with no mask and no sinks, key length 8192 or more, equal query and value head dims, and GQA 8 with head dim 64 or 128, or GQA 12 or 16 with head dim 128. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- The GQA kernel lets each simdgroup own a token sub-chunk and compute HPT of its group's query heads from registers, so each K/V byte is read `gqa_factor / HPT` times instead of `gqa_factor`; HPT is 4 for GQA 8 and 12 at head dim 128 and 2 for GQA 16, bounded by the 32 KB threadgroup memory cap (GQA 16 with HPT 4 would need 32.5 KB) and by register spilling at HPT 8 with head dim 128. — [source](https://github.com/ml-explore/mlx/pull/4380)
- For query length above 8 the full steel attention kernel runs. Without NAX its query tile is 32 rows and its key tile is 32 when the head dim is 256, the dtype is not float32 and the class is 'd', or when the head dim is below 128, and 16 otherwise, with 4 simdgroups along M and 2 along the head dim at head dim 256. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- With NAX (generation 17 or later, 18 or later for class 'p', macOS 26.2 or later) the full kernel uses a 64-row query tile (32 at head dim 512), key tile 32, 4 simdgroups along M (2 at head dim 512) and splits the head dim across `head_dim / 128` simdgroups at head dims 256 and 512. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- The vector kernels accept query length times GQA factor of at most 32 and query length of at most 8, so the vector-to-full switch happens at `min(8, floor(32 / gqa))` query rows: 8 for GQA 4, 6 for GQA 5, 4 for GQA 8. — source: `asserted`
- 2026-02-05: PR 3099 "Fix 2pass sdpa on < M2" (awni) found that on M1 and M2 with bfloat16, `blocks = 128` as a function constant lets the kernel run fewer than 1024 threads per threadgroup while 64 and 256 do not; it closed mlx-lm issue 844 (garbage output after about 1000 tokens), and removing the `blocks` function constant from pass 1 entirely caused a clear regression. — [source](https://github.com/ml-explore/mlx/pull/3099)
- 2026-05-11: PR 3455 added `MLX_SDPA_BLOCKS`; a 2-rank M4 Ultra cluster at 50k keys on a 256-expert MoE gained 6.5% decode at `blocks = 88` against the heuristic's 1024, with a sharp cliff from 92, which the author matched to about 352 concurrent simdgroups (4 KV heads times 88). — [source](https://github.com/ml-explore/mlx/pull/3455)
- 2026-06-28: PR 3637 added the (192, 128) asymmetric head-dim vector kernel; MiMo-V2.5 decode on an M3 Ultra measured 1.41x at 256 keys and 2.32x at 32768 keys (599 us fused against 1392 us fallback). — [source](https://github.com/ml-explore/mlx/pull/3637)
- 2026-07-20 (opened; shipped in 0.32.1): PR 3875 rounds `MLX_SDPA_BLOCKS` up to a multiple of 32; before it, values 16, 31, 33, 48 and 100 gave maximum absolute errors of 7.4e-2 to 1.2e-2 on an M5 at 8192 keys against 2.5e-5 for the default. — [source](https://github.com/ml-explore/mlx/pull/3875)
- 2026-08-05: PR 4018 added the missing threadgroup-size check to the 1-pass dispatch. — [source](https://github.com/ml-explore/mlx/pull/4018)
- 2026-08-18: PR 4077 added the GQA-8 pass-1 kernel; 2026-08-27 PR 4380 extended it to GQA 12 and 16 (in 0.32.3); 2026-09-02 PR 4431 added the batch offset it was missing (in 0.32.3). — [source](https://github.com/ml-explore/mlx/releases/tag/v0.32.3)
- 2026-10-03: PR 4596 unrolls the key loop of the generic 2-pass pass 1 by 4. — [source](https://github.com/ml-explore/mlx/pulls?q=is%3Apr+sdpa+2pass)
- 0.32.1 also lists PR 3843 (`unroll_count(4)` in the NAX attention Q@K.T loop) and PR 4330 (head dim 72 in Metal full attention). — [source](https://github.com/ml-explore/mlx/releases/tag/v0.32.1)
- An over-limit threadgroup dispatch is dropped by the driver with no error and the output buffer stays zeros; on an M1 Max the head-dim-512 pipelines cap at 832 threads where head-dim-256 allows 1024, and only the 2-pass dispatch had a check until PR 4018. — [source](https://github.com/ml-explore/mlx/pull/4018)
- Issue 4022 (M3 Max `applegpu_g15s`, Qwen3-VL-8B 4-bit, GQA 32/8, head dim 128): output was correct for 119 steps and became 904 `!` tokens from the first step at key length 1024, where routing flips to 2-pass with `blocks = 64`; forcing 1-pass fixed it, and `sdpa_vector.h` is byte-identical between v0.31.1 and v0.32.0. — [source](https://github.com/ml-explore/mlx/issues/4022)
- A second M3 Max owner could not reproduce it in pure MLX (keys 512 to 8192, bf16 and fp16, sliced KV cache, minus-infinity masks, v0.31.1 and v0.32.0), the reporter could not break it on an M5 Air, and zcbenz closed the issue on 2026-08-16 for lack of activity; the cause is unresolved. — [source](https://github.com/ml-explore/mlx/issues/4022)
- The GQA kernel's missing batch term made batched decode (B=2, 8192 keys) fail where B=1 was byte-identical. — [source](https://github.com/ml-explore/mlx/pull/4431)
- Since PR 3875 the override cannot express 88: it rounds to 96, which is above the reported cliff at 92 for that workload. — source: `asserted`
- For query length of 8 or less with query length times GQA above 32, `has_fused_kernel` returns false and the call decomposes into the unfused path; `force_fused` makes `use_fallback` raise instead of decomposing. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- Issue 3826 measured the cliff at 8 to 12 rows on an M5 Max (GQA 32/8, fp16, head dim 128): 0.813 ms to 1.906 ms at 16384 keys and 1.319 ms to 3.633 ms at 32768, a flat plateau to 48 rows and recovery at 64 (1.588 ms); the PR 3838 author measured the same step on an M3 Max at 8 to 9 rows (1.11 ms to 2.38 ms). — [source](https://github.com/ml-explore/mlx/issues/3826)
- Kernel author against maintainer on a multi-row 2-pass kernel for 8 to 16 rows (PR 3838): measured 1.92x at 9 rows and 1.23-1.33x at 16 rows on M3 Max (1.95-2.19x at 12 and 1.56-1.72x at 16 on M5 Max, confirmed by a second user on a source build with NAX verified active). zcbenz closed it on 2026-08-06 because a new 2-pass kernel has review and maintenance cost and is unlikely to be used in production. A commenter on the PR said MTP is rarely effective when drafting 8 or more tokens. — [source](https://github.com/ml-explore/mlx/pull/3838)
- Two reports of the same M3 Max class diverge: issue 4022 shows garbage at 1024 keys on 2-pass; a second M3 Max could not reproduce, so the failure may depend on the build or toolchain rather than the chip. — [source](https://github.com/ml-explore/mlx/issues/4022)
- Why M3 Max `applegpu_g15s` fails in issue 4022 while another M3 Max does not. — source: `asserted`
- Whether the 's' and 'd' block tables were tuned on measurements per chip generation (M5 Pro and Max both report class 's' with generation 17) or inherited from M1-M3; no source states it. — source: `asserted`
- Whether `MLX_SDPA_BLOCKS` between 32-multiples would help; the rounding forbids testing it. — source: `asserted`
- MLX picks Metal attention tiles and block counts from the last character of the reported GPU architecture ('p', 'g', 's', 'd'). — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/device.cpp)
- The 2-pass vector kernel is used for class 'd' or 's' at 1024 or more keys, or for GQA at 4096 or more keys on any class. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- An M5 MacBook Air switches to 2-pass at 4096 keys where an M3 Max switches at 1024. — [source](https://github.com/ml-explore/mlx/issues/4022)
- The 2-pass block count for class 's' runs 64, 128, 256, 512, 1024 as key length passes 1024, 8192, 32768, 65536 when `n_simds` is above 4. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- The 2-pass block count for class 'd' starts at 128 and reaches 256, 512 or 1024 by `n_simds` and key length. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- Other classes that reach 2-pass use 64 blocks for `n_simds` of 4 or more and 32 otherwise. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- `MLX_SDPA_BLOCKS` values are rounded up to a multiple of 32 and the environment-variable page documents that rounding. — [source](https://ml-explore.github.io/mlx/build/html/usage/environment_variables.html)
- `MLX_SDPA_BLOCKS` values 16, 31, 33, 48 and 100 gave errors from 1.2e-2 to 7.4e-2 before PR 3875 where the default gave 2.5e-5. — [source](https://github.com/ml-explore/mlx/pull/3875)
- At `blocks = 88` a 2-rank M4 Ultra cluster gained 6.5% decode at 50k keys on a 256-expert MoE, with a cliff from 92. — [source](https://github.com/ml-explore/mlx/pull/3455)
- PR 3099 found that `blocks = 128` as a function constant breaks the thread-count limit on M1 and M2 for bfloat16 and fixed mlx-lm issue 844. — [source](https://github.com/ml-explore/mlx/pull/3099)
- The GQA pass-1 kernel needs no mask, no sinks, one query row, 8192 or more keys, and GQA 8 (head dim 64 or 128) or GQA 12 or 16 (head dim 128). — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- On an M5 Pro the GQA-8 kernel gained 1.07-1.14x at head dim 128 and 1.15-1.27x at 64q/8kv/head dim 64 (8K to 32K keys), and 3.5% to 9.5% end-to-end decode on a truncated Qwen3-30B-A3B. — [source](https://github.com/ml-explore/mlx/pull/4077)
- The GQA 12 and 16 extension measured 1.14-1.18x and 1.07-1.11x on an M5 Pro against a 1.00-1.02x byte-identical control. — [source](https://github.com/ml-explore/mlx/pull/4380)
- GQA-kernel HPT is capped by 32 KB of threadgroup memory (GQA 16 with HPT 4 needs 32.5 KB) and by register spills at HPT 8. — [source](https://github.com/ml-explore/mlx/pull/4380)
- Release 0.32.3 lists PRs 4380 and 4431, and 0.32.1 lists PRs 3843, 3875 and 4330. — [source](https://github.com/ml-explore/mlx/releases/tag/v0.32.3)
- PR 4596 (merged 2026-10-03) unrolls the generic 2-pass key loop by 4 and measured 1.12x to 1.66x for GQA 1 and 1.29x to 1.49x for GQA 4 on an M1 Pro (`applegpu_g13s`). — [source](https://github.com/ml-explore/mlx/pull/4596)
- The no-NAX full attention kernel uses a 32-row query tile and a key tile of 16 or 32 chosen by head dim, dtype and class 'd'. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- The NAX full attention kernel uses a 64-row query tile (32 at head dim 512), a key tile of 32 and a head-dim split at 256 and 512. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/scaled_dot_product_attention.cpp)
- An over-limit threadgroup dispatch is silently dropped and returns zeros; M1 Max head-dim-512 pipelines cap at 832 threads. — [source](https://github.com/ml-explore/mlx/pull/4018)
- Issue 4022 reports garbage decode on an M3 Max at 1024 keys that forcing 1-pass fixed, that a second M3 Max could not reproduce, and that was closed for inactivity on 2026-08-16. — [source](https://github.com/ml-explore/mlx/issues/4022)
- Issue 3826 measured a 2.3x to 2.75x latency step from 8 to 12 query rows on an M5 Max, flat to 48 rows. — [source](https://github.com/ml-explore/mlx/issues/3826)
- PR 3838's multi-row 2-pass kernel gained 1.56-2.19x on an M5 Max at 12 and 16 rows and was closed unmerged on 2026-08-06. — [source](https://github.com/ml-explore/mlx/pull/3838)
- PR 3637's (192, 128) vector kernel measured 1.29x to 2.32x over the fallback on an M3 Ultra. — [source](https://github.com/ml-explore/mlx/pull/3637)
- The vector-to-full switch row is `min(8, floor(32 / gqa))`, so it moves with the GQA factor (8 rows at GQA 4, 4 rows at GQA 8). — source: `asserted`
- On 'g' and 'p' class chips multi-head attention never uses 2-pass, and an Ultra uses at least 128 blocks where a Max starts at 64. — source: `asserted`

## Corrections and disagreements

- CONTRADICTS (partly): metal-attention-multi-row-verify-cliff-at-6-15-q.md says no source explains why the cliff starts near 6 rows and that past rows times GQA above 32 attention falls to an unfused path. The source shows two regimes: for query length 8 or less the limit is rows times GQA at most 32 (so the cliff row is `floor(32 / gqa)`, and it moves with the model's GQA factor), and for query length above 8 supported head dims (64, 72, 80, 96, 128) go to the fused full kernel, not the unfused path. A 6-to-15 range is what several GQA factors produce together. — source: `asserted`
