Skip to content

Gemma 4: D512 two-pass SDPA + GQA dedup kernels - #13

Open
davidtai wants to merge 11 commits into
mainfrom
feat/gemma4-mlx-perf-stack
Open

davidtai wants to merge 11 commits into
mainfrom
feat/gemma4-mlx-perf-stack

Conversation

@davidtai

@davidtai davidtai commented Sep 5, 2026 •

Copy link
Copy Markdown

Gemma 4: D512 two-pass SDPA + GQA dedup kernels

This is the base (kernel) member of a four-repo stack for Gemma 4 26B-A4B. It adds
the MLX kernels the measured decode path needs. mlx-swift pins this commit.

1. What this PR changes

  • Adds a D=512 two-pass vector SDPA kernel for the five global layers. The stock
    vector SDPA does not reach the measured decode rate at head dim 512.
  • Adds a GQA dedup path on top of that kernel. It removes repeated key/value reads
    per query group.
  • Adds a softmax and SDPA-vector unroll. It cuts loop overhead in the decode
    attention.
  • Adds composed-prefill causal-bias synthesis. It builds the causal mask in the
    kernel instead of materializing it.
  • Adds NAX gather-QMM RHS elision. It drops a redundant right-hand-side copy in the
    expert gather.
  • Adds the affine-QMV tier family. It selects the matrix-vector tier by shape.

Each item is a default-on path. The C++ env gate =0 restores the stock path.

2. Why

The single largest decode lever is the D=512 two-pass SDPA
(DARKBLOOM_GEMMA4_D512_DECODE_2PASS, +9.4%; with _DEDUP, +10.4%). It is an MLX
kernel reached through MLXFast.scaledDotProductAttention. It is not an engine
Swift switch. So the engine cannot reach the measured decode on stock MLX.

Base port serial decode, one prompt, 17,408 tokens, 1,024 output, batch 1, M5 Max:

Arm Prefill tok/s Decode tok/s
Control (production pins) — 92.49
Base port (this stack, kernels on) 4,949 114.95

3. How it was tested

  • The test prompt is one real prompt of 17,408 tokens with 1,024 output tokens.
  • The batch size is 1 and decode is greedy.
  • Each case holds 40 °C and a quiet host before it runs.
  • Output is bit-exact to the stock kernels for every kernel except the D512 SDPA.
  • The D512 SDPA is a numeric near-tie (one greedy flip per ~6,000 tokens).

4. What is intentionally NOT changed

  • No batch-8 or ragged kernel is added. The serial path runs at batch 1 only.
  • Every added kernel keeps a =0 env gate to the stock path.

5. Stack

Merge bottom-up: mlx -> mlx-swift -> mlx-swift-lm -> d-inference.

# Repo Branch Tip PR
1 Layr-Labs/mlx feat/gemma4-mlx-perf-stack bb794a7 #13
2 Layr-Labs/mlx-swift feat/gemma4-mlx-perf eae562a Layr-Labs/mlx-swift#19
3 Layr-Labs/mlx-swift-lm feat/gemma4-mtp-stacked 0cb4c68d Layr-Labs/mlx-swift-lm#138
4 Layr-Labs/d-inference feat/gemma4-mtp-stacked 01374e6c Layr-Labs/d-inference#839

Pin dependencies:

  • mlx-swift pins this repo at the nested gitlink Source/Cmlx/mlx = bb794a7.
  • d-inference pins mlx-swift eae562a and mlx-swift-lm 0cb4c68d.

6. Known limits

  • The D512 SDPA is a near-tie of greedy output, not bit-exact. One greedy flip
    falls per ~6,000 tokens. HumanEval and MBPP gate this class (see the mlx-swift-lm
    member).

🤖 Generated with Claude Code

https://claude.ai/code/session_016mnuocRN7JMaSmAjSWBRWw

Refactor Pass

A dedicated refactor subagent reviewed the diff against the Layr-Labs AGENTS.md structure, clean-code and refactor rules (d-inference ml-explore#1338). Outcome: no change. The kernel code is review-only. The host code in scaled_dot_product_attention.cpp has clean boundaries and flat admission logic, with no dead code or duplication.

Two meaning conflicts with main abd33ee. The text merge is clean, so neither shows up as a conflict:

  1. Preserve FP32 partial sums in two-pass vector SDPA #16 breaks the default-on D512 decode path after a merge.

    • Preserve FP32 partial sums in two-pass vector SDPA #16 renames the 2-pass kernels to sdpa_vector_2pass_fp32partials_*, and the host now builds "sdpa_vector_2pass_fp32partials_1".
    • This PR's new instantiate_sdpa_vector_2pass macro still emits "sdpa_vector_2pass_1_" #type ….
    • In the merged tree, every D=512 non-dedup decode call would look up a kernel that does not exist. DARKBLOOM_GEMMA4_D512_DECODE_2PASS is on by default.
    • Fix at rebase: rename that macro's string to sdpa_vector_2pass_fp32partials_1_.
  2. Expose coherent allocator accounting and correct quantized bias sums #15 changes load_vector to sum in U, which makes two claims in this PR false:

    The author must decide: follow Expose coherent allocator accounting and correct quantized bias sums #15 with U(...) casts and measure again, or keep the old form and correct the comments.

For the author:

  • The host comments carry work-wave labels and long narrative. At rebase, cut them to their invariants: D512 always takes the 2-pass form, empty blocks fold in with weight 0, and dedup stays off until the register load is measured.
  • scaled_dot_product_attention.metal:50 cites a stale line number (:461; the launch is now near :474).
  • No tests cover D=512 SDPA, the gqa SPLIT=2 path or the new QMV tiers. A weight-free test_fast_sdpa.py case at D=512 is possible; the dedup path needs 8192 or more keys.
  • The only check on this head is CodeQL. The fork build/test workflow has not run on it.

Port and fix (follow-up):

Refactor pass on the fix, e3905de: a dedicated refactor subagent shortened the D512 comment block in scaled_dot_product_attention.metal to 2 lines (the comment only; the D512-2PASS tag stays), and set the test tolerances and the dedup scale once. All assertions are unchanged. black and isort pass, and test_fast_sdpa.py passes (27). For the author: instantiate_sdpa_vector_2pass repeats upstream's 2-pass instantiate_kernel call, and that copy is how the name drifted. Sharing one macro would edit upstream code.

U-cast fix, e602e61 (follow-up, David's ruling):

  • mma8_runsum4 and qmv_fast_singlerow_affine2_g64 now sum the 4-tuple in fp32, the same as Expose coherent allocator accounting and correct quantized bias sums #15's load_vector, and the comments are true again. A new test, test_qmv_bias_sum_widens_inputs, fails before the fix (65536 where 66560 is expected) and passes after it.
  • Error against the post-Expose coherent allocator accounting and correct quantized bias sums #15 reference on gated outputs: max rel error was 4.8–10.5 before and is now ≤ 1 bf16 ulp. Argmax flips were 5–12 per 400 rows on mma8 bf16 and are now 0. Against an fp64 ground truth, each tier's error now equals the reference's.
  • test_quantized.py passes 39/39 (with MLX_ENABLE_TF32=0), and test_nn.py passes 73. check_forkdiff, black and isort pass.
  • Findings for the author:
    1. Upstream qmv_wide (Add small-batch quantized matvec kernel (qmv_wide) ml-explore/mlx#3764) sends every affine call with M ≥ 2 to qmv_wide on GPU generation 15 or newer (quantized.cpp:543 and :1987). So the mma8, quad_stream and pair tiers never run on M3 or newer, including the M5 and the M4 Pro ranked box.
    2. The singlerow gate checks only out_vec_size == 98336, M=1 and 2-bit, not K. Its loop steps 1024 values while the host requires only K % 512 == 0, so K=1536 could read past the row.

davidtai and others added 6 commits September 3, 2026 10:30
Chunk M1 of the Gemma 4 mlxfast port (ledger §C.1).
Ported from Layr-Labs/mlxfast-gemma4-26b-a4b-engine 944c572 (final
vendored tree), which carried it from engine commits 796aa221 (Validate
submission 6ce1e46e-fe7e-4996-be64-cfb9501fc8f5, softmax) and 57087ca2
(Accept submission 82c69b6c-37e9-4f38-9d20-6b122e7ceb57, sdpa_vector).

Mechanism: `#pragma unroll` on the compile-time-trip-count loops —
N_READS in both softmax kernels, qk_per_thread / v_per_thread /
elem_per_thread in the SDPA vector, vector-2pass and vector-2pass-reduce
kernels. No arithmetic, no accumulation order, no operand shapes change;
these are loop-structure hints only, so every output stays bit-identical.

Files:
- mlx/backend/metal/kernels/softmax.h (+12)
- mlx/backend/metal/kernels/sdpa_vector.h (+14)

Already upstream in 0.32.2: nothing. Both hunks are new here. `softmax.h`
is byte-identical between the engine's fork base (d5a2404) and this
branch's tip, so it applied unchanged. `sdpa_vector.h` drifted 143 lines
across 0.32.0 -> 0.32.2, but every pragma landed by three-way merge
against d5a2404 with no conflict.

Co-authored-by: fkiene <[email protected]>
Co-authored-by: jungjipdo <[email protected]>

Co-Authored-By: Claude Fable 5.1 <[email protected]>
Claude-Session: https://claude.ai/code/session_01NS86aXV85ai4eLRw1ADEyX
… fused GEMM epilogues

Chunk M2 of the Gemma 4 mlxfast port (ledger §C.1).
Ported from Layr-Labs/mlxfast-gemma4-26b-a4b-engine 944c572 (final
vendored tree), which carried CAUSAL-CLOAD from engine commit 4f44957c
(Validate submission 5a454a6a-26b6-4af0-ab20-d256cfe328bd) and
NAX-SKIP-EMPTY-001 from f68023e0 (Validate submission
1e9f5531-235a-4920-b516-c08a9908d864).

Mechanism (CAUSAL-CLOAD): in the addmm epilogue, recognise the
composed-prefill causal-bias operand by its signature — bf16 accumulate,
!transpose_a && transpose_b, fdc == 1, ldc == N + 1, M <= N, all-zero C
batch strides — and synthesize its two constants per accumulator element
instead of loading them: widened bfloat16 lowest finite (0xFF7F) strictly
above the causal diagonal at N - M, widened bfloat16 negative zero on and
below it. A row stride of N + 1 cannot arise from a contiguous or
broadcast operand of the declared output width, so the signature is
unambiguous. The addend still enters through the same TransformAdd with
the same widening as the loaded operand it replaces, so every stored word
is bit-identical; every other addmm keeps the loaded-operand epilogue.
Adds the `kCausalBiasSynthEligible` constexpr gate (so complex64 never
instantiates the branch), the `c_bstride_zero` batch-stride check, and the
`gemm_epilogue_causal_synth` NAX tile helper.

Files:
- mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h (+57/-2)
- mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h (+72/-2)

Already upstream in MLX 0.32.2, so NOT re-applied:
- NAX-SKIP-EMPTY-001 in its entirety. This branch's tip already carries the
  empty-output-extent elision unconditionally — `if constexpr (!kAlignedM ||
  !kAlignedN) { if (!has_output) …}` in steel/gemm/gemm_nax.h and the
  `(kAlignedM.value || sgp_sm > 0) && (kAlignedN.value || sgp_sn > 0)`
  epilogue guard in steel_gemm_fused_nax.h. The engine's version of the
  same optimisation wraps those guards in a `DARKBLOOM_GEMMA4_NAX_SKIP_EMPTY`
  kill-switch macro; re-applying it would only make an already-live
  optimisation conditional. steel/gemm/gemm_nax.h is therefore untouched by
  this commit, and the macro definition and both guard rewrites were dropped
  from steel_gemm_fused_nax.h — only the CAUSAL-CLOAD body was kept inside
  upstream's guard.

Conflicts resolved (three-way against the engine's fork base d5a2404):
- steel/gemm/gemm_nax.h ×3 — kept this branch's unguarded skip-empty.
- steel_gemm_fused_nax.h ×1 — kept this branch's epilogue guard, took the
  engine's causal-synth body inside it.
steel_gemm_fused.h merged with no conflict (0 lines of drift vs d5a2404).

Co-authored-by: Amal-David <[email protected]>

Co-Authored-By: Claude Fable 5.1 <[email protected]>
Claude-Session: https://claude.ai/code/session_01NS86aXV85ai4eLRw1ADEyX
…kernel

Chunk M3 of the Gemma 4 mlxfast port (ledger §C.1).
Ported from Layr-Labs/mlxfast-gemma4-26b-a4b-engine 944c572 (final
vendored tree), aggregating engine commits 465ce5ce, 58d03e6a, 3b51325b
and f68023e0 (Accept/Validate submissions 91c92a7e-fba5-4fe7-abdc-a0e61d3e2a87,
6bf82ab1-f553-41de-a1d6-4041e4a0a352, c8726cfe-26bb-4ef9-a089-099d5243a5c8
and 1e9f5531-235a-4920-b516-c08a9908d864).

Mechanism: `affine_gather_qmm_rhs_nax` computes a full SM x SN tile per
simdgroup and then discards, at `store_slice`, every row outside the
current expert segment `[seg_lo, seg_hi)`. This hoists that band ahead of
the K-loop and skips the discarded work instead of computing it:
- `seg_empty` — the whole band is dead, so the simdgroup runs no A load
  and no MMA at all (band granularity);
- `seg_partial` (aligned-M only) — 16-row fragment-row granularity: only
  fragment rows intersecting the band call `gather_rhs_load_frag_row` /
  `gather_rhs_mma_frag_row`, each running the stock path's exact op
  sequence for that row.
Cooperative weight loads and every `threadgroup_barrier` stay
unconditional, so barrier convergence is preserved; `offset`/`offset_next`
are threadgroup uniform and `seg_*` simdgroup uniform, so no intra-simdgroup
divergence is introduced. Gated by `kGatherRhsSegmentElide`; with it off
only the stock path runs. Also brings the engine's `qmm_t_nax_tgp_impl`
and `tile_matmad_nax` additions in the same file.

Files:
- mlx/backend/metal/kernels/quantized_nax.h (+271/-34)

Already upstream in MLX 0.32.2, and therefore SUPERSEDED rather than
re-applied: this branch's tip had independently added the band-granular
half of the same elision as `sg_active`, computed from `m_lo_lim`/
`m_hi_lim` — expressions textually identical to the engine's `seg_lo`/
`seg_hi`. `seg_empty` is exactly `!sg_active`, and `seg_partial` is the
finer tier upstream does not have, so the engine's form subsumes it. The
now-unreferenced `m_lo_lim`/`m_hi_lim`/`sg_active` trio was removed.

Conflicts resolved (three-way against the engine's fork base d5a2404), all
five inside `affine_gather_qmm_rhs_nax`:
- K-loop head ×2 and unaligned-K tail ×1 — took the engine's `seg_partial`/
  `seg_empty` structure over this branch's `if (sg_active)`.
- Btile load reformat ×2 — pure whitespace; took the engine's wrapping.
- store block ×1 — took the engine's `if (!seg_empty)` + `seg_lo`/`seg_hi`
  spelling of this branch's `m_lo_lim`/`m_hi_lim` slice.

NOT VALIDATED ON DEVICE. This port was produced build-only; the numerics
of the fragment-row path have not been re-measured against this branch's
kernels.

Co-authored-by: Amal-David <[email protected]>
Co-authored-by: i34-9 <[email protected]>

Co-Authored-By: Claude Fable 5.1 <[email protected]>
Claude-Session: https://claude.ai/code/session_01NS86aXV85ai4eLRw1ADEyX
Chunk M4 of the Gemma 4 mlxfast port (ledger §C.1); M4a and M4b are
combined here — see "Deviation" below.
Ported from Layr-Labs/mlxfast-gemma4-26b-a4b-engine 944c572 (final
vendored tree), aggregating engine commits 61c33ed1, fd981eea, eb82740f,
3ab9bd49, 41ab2004, 46e41ab9, 2afa9b80, 1356b788, 15b14f54, cbbcb92d,
bdbb9947, 38269aa8, 6042598b, 712602f4, 81314800, cdddcf15, 60318626,
f6904ac8, c2284eeb and b0631d8e (Accept/Validate submissions on the
ranked Gemma 4 26B-A4B track).

Mechanism — new tiers, all bit-exact restatements of the stock `qdot` /
`qmv_impl` arithmetic with a different load or dispatch shape:
- `qdot_affine4_registered{,_word}`, `qdot_affine8_registered{,_word}`,
  `qdot_affine4_pair{,_word}`, `qdot_affine8_pair`, `qdot_affine4_loaded`,
  `qdot_affine4_loaded_pair`, `qdot_affine4_g64_word` — the bits==4 /
  values_per_thread==8 (and byte-weight) arms of `qdot` against packed
  weight words already in registers: same nibble masks, same two 4-term
  sums, same accumulation order over i, same `scale * accum + sum * bias`
  close. Only the load shape differs (one 4-byte load instead of two
  2-byte loads).
- `qmv_affine4_g64_pair_impl`, `_triple_stream_impl`, `_quad_stream_impl`,
  `qmv_affine8_g64_pair_impl`, `_quad_stream_impl`, and
  `qmv_affine4_g64_singles_impl` — 1/2/3/4 same-expert assignments served
  from ONE weight stream; each (output, input) pair keeps its own
  accumulator and K-loop order, so every output element's add sequence
  matches the incumbent per-arm kernel.
- `qmv_fast_crossrow_affine4_g64{,_wide,_m}`, `qmv_fast_singlerow_affine2_g64`
  — cross-row tight-grid bodies for the batch-8 decode plane.
- `mma8_lane`/`mma8_lo`/`mma8_hi`/`mma8_runsum4` +
  `gemma4_qmv_mma8_affine4_g64_impl` — fp32 `simdgroup_float8x8` body for
  the M=8 decode cohort on 4-bit affine g64 (A = raw weight codes 8x8,
  B = x-transpose 8x8, C zeroed per g64 group).
- `gather_qmv_gemma4_down_tile` + the `affine_gather_qmv` dispatch rewrite
  — RUN-QUAD leader election over the flattened 64-assignment route table,
  reading the EXPERT-PREFIX-BOUNDS-001 packed route word (bit 31 = format
  flag, bits 0-7 expert, 8-13 run offset, 14-19 run length) with a
  linear-scan fallback when the flag is clear, plus the y-tile-coarsened
  arm for the K = 704 down plane. Both arms are compile-time flippable
  (`gemma4_down_tile`) and bit-identical by construction.
- Two `qmv_impl` loop bounds change from `k < in_vec_size - block_size` to
  `k <= …`, so an exactly block-aligned input runs its last full block on
  the fast path instead of the `qdot_safe` tail; the tail's `remaining`
  clamp already covers k == in_vec_size.

Files:
- mlx/backend/metal/kernels/quantized.h (+2281/-55)

Already upstream in MLX 0.32.2 and preserved unchanged by the three-way
merge: this branch's `qmv_wide` family (Layr-Labs/mlx-swift 606d28c
"expose qmv_wide to Swift runtime", 4 references) and the
`has_global_scale` template parameter added across the affine kernels
(10 references) both live in the same `qmv_affine*` region the engine's
tiers were written into. Neither was reverted; the engine's fork base
(d5a2404) predates both.

Conflict resolved (one, three-way against d5a2404): the declaration
immediately preceding `[[kernel]] void affine_gather_qmv` — the engine
inserted its `qdot_affine4_g64_word` + `qmv_affine4_g64_singles_impl` +
`gather_qmv_gemma4_down_tile` block there while this branch had widened
the following template to `template <typename T, int group_size, int bits,
bool has_global_scale = false>`. Kept this branch's widened signature and
inserted the engine's block ahead of it.

Deviation from the planned chunking: the ledger suggests splitting this
into M4a (the `qdot`/`_pair`/`_stream` primitives) and M4b (the
Gemma-4-specific tiers). The diff does not split at hunk boundaries — one
888-line hunk contains both `qmv_affine4_g64_pair_impl` and the
`mma8_*`/`gemma4_qmv_mma8_*` family — so a split would have required
sub-hunk surgery on generated kernel text with no device validation
available. Kept as one commit.

NOT VALIDATED ON DEVICE. Build-only port; none of these tiers has been
re-measured against this branch's kernels.

Co-authored-by: 0xkydo <[email protected]>
Co-authored-by: Amal-David <[email protected]>
Co-authored-by: DashiellB <[email protected]>
Co-authored-by: brandonegg <[email protected]>
Co-authored-by: delordemm1 <[email protected]>
Co-authored-by: ercumentyildirim <[email protected]>
Co-authored-by: exakoss <[email protected]>
Co-authored-by: i34-9 <[email protected]>
Co-authored-by: ivanfioravanti <[email protected]>
Co-authored-by: jungjipdo <[email protected]>
Co-authored-by: newjordan <[email protected]>
Co-authored-by: polymorf <[email protected]>
Co-authored-by: rinaldofesta <[email protected]>
Co-authored-by: rube-de <[email protected]>
Co-authored-by: samfenwick <[email protected]>

Co-Authored-By: Claude Fable 5.1 <[email protected]>
Claude-Session: https://claude.ai/code/session_01NS86aXV85ai4eLRw1ADEyX
C1 rung 2. Gemma 4 26B-A4B's five global attention layers are 16 query
heads / 2 KV heads at head_dim 512. `has_fused_kernel`'s vector head-dim
list is {64, 96, 128, 192, 256}, so every decode call on those layers is
rejected and `use_fallback` sends it to the unfused
`scale.q -> matmul -> softmax -> matmul` graph (fast.cpp). That graph
unflattens the query to [B, kv_heads, gqa, 1, D] and broadcasts K and V
over the gqa axis, so its two matmuls are batched gemvs over 16 batch
entries against 2 distinct planes: each key plane and each value plane is
streamed once PER QUERY HEAD -- eight times per layer -- and a
[B, 16, 1, kL] bf16 score plane is materialised, written once and read
twice.

`sdpa_vector_2pass_1` and `sdpa_vector_2pass_2` are already templated on
D and V, so no new kernel body is needed: this instantiates them at
512/512 and admits the dim.

* kernels/scaled_dot_product_attention.metal: a 2-pass-only instantiation
  macro plus `..._2pass(type, 512, 512)` and
  `..._aggregation(type, 512)`. The single-pass `sdpa_vector` twin is
  deliberately NOT instantiated -- it holds q, k and o at D/32 floats each
  and launches at 1024 threads per threadgroup, which at D = 512 is 48
  live floats against the Metal maximum thread count, an occupancy claim
  the split-K kernel (32 x gqa_factor threads, 32 live floats) does not
  make.
* scaled_dot_product_attention.cpp: `eval_gpu` routes every D = 512
  vector call to `sdpa_vector_2pass` at ANY key length, since there is no
  single-pass instantiation to fall back to. The 2-pass form is
  length-generic: blocks that see no key leave `sums = 0` and
  `maxs = finite_min`, which the merge pass folds in with weight
  `exp(finite_min - max) == 0`.

Prefill is untouched: at query length > 8 the call takes the full
attention branch, whose head-dim list is unchanged, so it keeps falling
back exactly as before. MTP verify rectangles are also untouched --
`query_sequence_length * gqa_factor > 32` rejects them at gqa 8 and any
L > 4.

NOT bit-exact against the unfused graph: the split-K kernel carries an
online (running-max) softmax and folds `blocks` partials in a second
pass, so the reduction order over the key axis differs. The bar is
greedy-token parity.

Switch DARKBLOOM_GEMMA4_D512_DECODE_2PASS, default ON; any of
{0, false, no, off} removes the admission and restores the unfused graph
byte for byte.

Also instantiates `sdpa_vector_2pass_1_gqa` at 512/HPT=2 and extends the
`_gqa` kernel-name condition to reach it, behind
DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP, DEFAULT OFF. The plain 2-pass
kernel gives each query head its own simdgroup and so still reads each
K/V byte gqa_factor times; the dedup variant reads it gqa_factor / HPT
times and is the only form that actually removes the redundant stream.
It is off by default because at D = 512 a thread holds
HPT * (D / 32) * 2 = 64 live floats at 256 threads per threadgroup, and a
pipeline whose `maxTotalThreadsPerThreadgroup` came back under 256 would
make `check_kernel_threadgroup_size` throw rather than degrade. Turn it
on once that is measured on the device.

Co-Authored-By: Claude Fable 5.1 <[email protected]>
Claude-Session: https://claude.ai/code/session_014L39jXT8ReGzjxUoUKLfan
…up limit

The DEDUP arm failed pipeline creation on the device:

  [metal::Device] Unable to load kernel
  sdpa_vector_2pass_1_gqa_bfloat16_t_512_512_nomask_qnt_nc_nosinks_128:
  Threadgroup memory size (32896) exceeds the maximum threadgroup memory
  allowed (32768)

Cause. `sdpa_vector_2pass_1_gqa` publishes its cross-simdgroup merge plane as
`threadgroup U o_sh[G * HPT * V]`, U = float. At (G 8, HPT 2, V 512) that is
8192 floats = 32,768 B on its own, and `se_sh` + `mx_sh` (16 floats each) put
it 128 B over the limit. Exactly the reported 32,896. `blocks` is not a term
in that expression, so tuning MLX_SDPA_BLOCKS could not have moved it; the
shipped (64, HPT 8) and (128, HPT 4) instantiations both land at 16,640 B,
which is why the limit had never been reached before.

Fix. A `SPLIT` template parameter: the plane is published in SPLIT passes of
V / SPLIT columns, so it allocates `G * HPT * V / SPLIT` floats. The 512
instantiation takes SPLIT = 2 -> 16,512 B, in line with the shipped two. The
existing instantiations take SPLIT = 1 and are unchanged.

SPLIT does not change the arithmetic:

* each lane keeps the same register slice, and the shared plane is only a
  scratch relabelling of that slice -- write and read use the identical lane
  mapping (`simd_lid * v_per_pass`), never the global column index, so the
  plane's internal layout never has to match the output column order;
* `gmax` and `denom` are computed once, on pass 0, from the full `mx_sh` /
  `se_sh` arrays, in the same order over s;
* every `acc[i]` keeps its accumulation order over s inside its pass, and no
  output element is touched by more than one pass;
* SPLIT = 1 is the shipped body instruction for instruction -- the publish of
  the plane and of the scalars stay in one loop, there is still exactly one
  barrier before the merge, and the extra write-after-read barrier is guarded
  by `p > 0`.

Also adds three `static_assert`s so this class of failure cannot reach a
device again: SPLIT must divide V and V / 32, and the threadgroup allocation
must fit 32,768 B. Verified both ways with the offline gate --
`xcrun metal -c` passes at SPLIT = 2 with no warnings, and temporarily
setting the 512 instantiation back to SPLIT = 1 reproduces the device failure
as a compile error naming `sdpa_vector_2pass_1_gqa<float, 512, 512, 8, 2, 1>`.

Symbol check on the resulting .air: `sdpa_vector_2pass_1_gqa_*_512_512`
present for all three types at <512, 512, 8, 2, 2>, and the shipped
64/128 kernels still at <..., 1>.

DEDUP stays DEFAULT OFF. The remaining unmeasured claim is registers: a
thread holds q[2][16] + o[2][16] + kr[16] + vr[16] + acc[16] = 112 live
floats at 256 threads per threadgroup, and a pipeline whose
`maxTotalThreadsPerThreadgroup` came back under 256 would throw from
`check_kernel_threadgroup_size`. Turn it on with
DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP=1 for its own arm.

Co-Authored-By: Claude Fable 5.1 <[email protected]>
davidtai and others added 5 commits October 4, 2026 19:45
Main #16 renamed the 2-pass vector SDPA kernels to
sdpa_vector_2pass_fp32partials_*. The host builds
"sdpa_vector_2pass_fp32partials_1_<type>_512_512", but the D512
instantiation still emitted "sdpa_vector_2pass_1_<type>_512_512", so
every head-dim 512 call failed with "Unable to load function".

Rename the instantiation. Replace the stale line reference to
scaled_dot_product_attention.cpp with the symbol name.

Add tests that compare with the reference attention: head dim 512
(16/2 and 4/4 heads, 7 to 8192 keys, with and without a mask), and the
D512 GQA dedup kernel (8192 and 8201 keys) in a child process with
DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP=1.

Co-Authored-By: Claude Opus 5.5 <[email protected]>
The fork-diff gate from main (#19) failed after the merge: softmax.h,
steel_gemm_fused.h and steel_gemm_fused_nax.h were not described by any
section. Add one section for this PR that lists every file it changes.

Co-Authored-By: Claude Opus 5.5 <[email protected]>
The D512-2PASS comment is now two lines, as AGENTS.md asks. The head-dim
512 test sets its tolerance once for each dtype, and the dedup child
script computes the scale once. No kernel code or assertion changes.

Co-Authored-By: Claude Opus 5.5 <[email protected]>


qmv_fast_singlerow_affine2_g64 and mma8_runsum4 added the 4-tuple of x
values in T before they widened the result. #15 changed load_vector to
widen each value to U first, so these two tiers no longer matched the
reference. Both now widen each value to float before the add, and the
comments that call them twins of load_vector are true again.

test_qmv_bias_sum_widens_inputs uses the bf16 tuple (256, 1, 1, 1),
where a sum in T loses 3 per tuple.

Co-Authored-By: Claude Opus 5.5 <[email protected]>
davidtai added a commit to Layr-Labs/mlx-swift that referenced this pull request Oct 5, 2026
Move Source/Cmlx/mlx from 9e239f9e8 to e602e6123 (Layr-Labs/mlx#13
head). Run tools/update-mlx.sh on the new pin.

- metal/quantized.h, quantized.cpp: mma8_runsum4 and
  qmv_fast_singlerow_affine2_g64 widen each value to float before the
  bias run sum, as load_vector does (mlx e602e6123).
- metal/scaled_dot_product_attention.metal: shorter D512-2PASS comment
  (mlx e3905de71).

Co-Authored-By: Claude Opus 5.5 <[email protected]>

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