From 7391caff44f2712552b1f43ac05073fae84eb025 Mon Sep 17 00:00:00 2001 From: David Tai <8346495+davidtai@users.noreply.github.com> Date: Thu, 3 Sep 2026 10:30:03 -0500 Subject: [PATCH 01/11] perf(metal): unroll the fixed-trip loops in softmax and SDPA-vector MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 <46886660+fkiene@users.noreply.github.com> Co-authored-by: jungjipdo <130676635+jungjipdo@users.noreply.github.com> Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01NS86aXV85ai4eLRw1ADEyX --- mlx/backend/metal/kernels/sdpa_vector.h | 14 ++++++++++++++ mlx/backend/metal/kernels/softmax.h | 12 ++++++++++++ 2 files changed, 26 insertions(+) diff --git a/mlx/backend/metal/kernels/sdpa_vector.h b/mlx/backend/metal/kernels/sdpa_vector.h index 3f40dbd7c0..23194a3f00 100644 --- a/mlx/backend/metal/kernels/sdpa_vector.h +++ b/mlx/backend/metal/kernels/sdpa_vector.h @@ -81,9 +81,11 @@ template out += o_offset * V + simd_gid * v_per_thread; // Read the query and 0 the output accumulator + #pragma unroll for (int i = 0; i < qk_per_thread; i++) { q[i] = static_cast(scale) * queries[i]; } + #pragma unroll for (int i = 0; i < v_per_thread; i++) { o[i] = 0; } @@ -107,12 +109,14 @@ template } if (use_key) { // Read the key + #pragma unroll for (int j = 0; j < qk_per_thread; j++) { k[j] = keys[j]; } // Compute the i-th score U score = 0; + #pragma unroll for (int j = 0; j < qk_per_thread; j++) { score += q[j] * k[j]; } @@ -130,6 +134,7 @@ template sum_exp_score = sum_exp_score * factor + exp_score; // Update the output accumulator + #pragma unroll for (int j = 0; j < v_per_thread; j++) { o[j] = o[j] * factor + exp_score * values[j]; } @@ -160,6 +165,7 @@ template sum_exp_score = simd_sum(sum_exp_scores[simd_lid] * factor); // Now we need to aggregate all the outputs + #pragma unroll for (int i = 0; i < v_per_thread; i++) { outputs[simd_lid * BD + simd_gid] = o[i]; threadgroup_barrier(mem_flags::mem_threadgroup); @@ -170,6 +176,7 @@ template // And write the output if (simd_lid == 0) { + #pragma unroll for (int i = 0; i < v_per_thread; i++) { out[i] = static_cast(o[i]); } @@ -248,6 +255,7 @@ template maxs += o_offset * blocks + block_idx; // Read the query + #pragma unroll for (int i = 0; i < qk_per_thread; i++) { q[i] = static_cast(scale) * queries[i]; } @@ -272,6 +280,7 @@ template if (use_key) { // Compute the i-th score U score = 0; + #pragma unroll for (int i = 0; i < qk_per_thread; i++) { score += q[i] * keys[i]; } @@ -290,6 +299,7 @@ template sum_exp_score = sum_exp_score * factor + exp_score; // Update the output accumulator + #pragma unroll for (int i = 0; i < v_per_thread; i++) { o[i] = o[i] * factor + exp_score * values[i]; } @@ -312,6 +322,7 @@ template maxs[0] = max_score; } + #pragma unroll for (int i = 0; i < v_per_thread; i++) { out[i] = static_cast(o[i]); } @@ -511,6 +522,7 @@ template U factor = fast::exp(maxs[simd_gid] - max_score); // Update the output accumulator + #pragma unroll for (int i = 0; i < elem_per_thread; i++) { o[i] += factor * static_cast(partials[i]); } @@ -520,6 +532,7 @@ template } // Use shared memory to transpose and reduce the final block + #pragma unroll for (int i = 0; i < elem_per_thread; i++) { outputs[simd_lid * BD + simd_gid] = o[i]; threadgroup_barrier(mem_flags::mem_threadgroup); @@ -530,6 +543,7 @@ template // And write the output if (simd_lid == 0) { + #pragma unroll for (int i = 0; i < elem_per_thread; i++) { out[i] = static_cast(o[i]); } diff --git a/mlx/backend/metal/kernels/softmax.h b/mlx/backend/metal/kernels/softmax.h index 6ea4ac7329..d995610649 100644 --- a/mlx/backend/metal/kernels/softmax.h +++ b/mlx/backend/metal/kernels/softmax.h @@ -27,10 +27,12 @@ template in += gid * size_t(axis_size) + lid * N_READS; if (lid * N_READS + N_READS <= axis_size) { + #pragma unroll for (int i = 0; i < N_READS; i++) { ld[i] = AccT(in[i]); } } else { + #pragma unroll for (int i = 0; i < N_READS; i++) { ld[i] = ((lid * N_READS + i) < axis_size) ? AccT(in[i]) : Limits::min; @@ -44,6 +46,7 @@ template // Get the max AccT maxval = Limits::finite_min; + #pragma unroll for (int i = 0; i < N_READS; i++) { maxval = (maxval < ld[i]) ? ld[i] : maxval; } @@ -63,6 +66,7 @@ template // Compute exp(x_i - maxval) and store the partial sums in local_normalizer AccT normalizer = 0; + #pragma unroll for (int i = 0; i < N_READS; i++) { AccT exp_x = softmax_exp(ld[i] - maxval); ld[i] = exp_x; @@ -85,10 +89,12 @@ template // Normalize and write to the output out += gid * size_t(axis_size) + lid * N_READS; if (lid * N_READS + N_READS <= axis_size) { + #pragma unroll for (int i = 0; i < N_READS; i++) { out[i] = T(ld[i] * normalizer); } } else { + #pragma unroll for (int i = 0; i < N_READS; i++) { if ((lid * N_READS + i) < axis_size) { out[i] = T(ld[i] * normalizer); @@ -123,20 +129,24 @@ template int offset = r * lsize * N_READS + lid * N_READS; AccT vals[N_READS]; if (offset + N_READS <= axis_size) { + #pragma unroll for (int i = 0; i < N_READS; i++) { vals[i] = AccT(in[offset + i]); } } else { + #pragma unroll for (int i = 0; i < N_READS; i++) { vals[i] = (offset + i < axis_size) ? AccT(in[offset + i]) : Limits::min; } } prevmax = maxval; + #pragma unroll for (int i = 0; i < N_READS; i++) { maxval = (maxval < vals[i]) ? vals[i] : maxval; } normalizer *= softmax_exp(prevmax - maxval); + #pragma unroll for (int i = 0; i < N_READS; i++) { normalizer += softmax_exp(vals[i] - maxval); } @@ -175,10 +185,12 @@ template r++) { int offset = r * lsize * N_READS + lid * N_READS; if (offset + N_READS <= axis_size) { + #pragma unroll for (int i = 0; i < N_READS; i++) { out[offset + i] = T(softmax_exp(in[offset + i] - maxval) * normalizer); } } else { + #pragma unroll for (int i = 0; i < N_READS; i++) { if (offset + i < axis_size) { out[offset + i] = From e2f1d5e2238108f3b60c15ac0e4e8b3bf11c06db Mon Sep 17 00:00:00 2001 From: David Tai <8346495+davidtai@users.noreply.github.com> Date: Thu, 3 Sep 2026 10:31:11 -0500 Subject: [PATCH 02/11] perf(metal): synthesize the composed-prefill causal bias in the steel fused GEMM epilogues MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 <11647194+Amal-David@users.noreply.github.com> Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01NS86aXV85ai4eLRw1ADEyX --- .../steel/gemm/kernels/steel_gemm_fused.h | 60 +++++++++++++++- .../steel/gemm/kernels/steel_gemm_fused_nax.h | 72 ++++++++++++++++++- 2 files changed, 128 insertions(+), 4 deletions(-) diff --git a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h index 85830872d1..46c4f35f11 100644 --- a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h +++ b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h @@ -70,6 +70,13 @@ template < return; } + // CAUSAL-CLOAD eligibility, template facts only. The synthesis below + // fires solely on the composed-prefill causal-bias signature; every other + // addmm keeps the loaded-operand epilogue untouched. + constexpr bool kCausalBiasSynthEligible = + !transpose_a && transpose_b && metal::is_same_v; + bool c_bstride_zero = true; + // Adjust for batch if (has_batch) { const constant auto* A_bstrides = batch_strides; @@ -84,6 +91,9 @@ template < if (use_out_source) { const constant auto* C_bstrides = B_bstrides + params->batch_ndim; C += elem_to_loc(tid.z, batch_shape, C_bstrides, params->batch_ndim); + for (int d = 0; d < params->batch_ndim; d++) { + c_bstride_zero = c_bstride_zero && (C_bstrides[d] == 0); + } } } else { A += params->batch_stride_a * tid.z; @@ -91,6 +101,7 @@ template < if (use_out_source) { C += addmm_params->batch_stride_c * tid.z; + c_bstride_zero = addmm_params->batch_stride_c == 0; } } @@ -195,8 +206,53 @@ template < mma_op.apply_epilogue( C, addmm_params->ldc, addmm_params->fdc, epilogue_op_axpby); } else { - mma_op.apply_epilogue( - C, addmm_params->ldc, addmm_params->fdc, epilogue_op_add); + // The synthesis touches BlockMMA members that only the real-typed + // specialization has; constexpr-gate it so ineligible element types + // (complex64) never instantiate the branch. + bool synthesized = false; + if constexpr (kCausalBiasSynthEligible) { + if (addmm_params->fdc == 1 && + addmm_params->ldc == params->N + 1 && params->M <= params->N && + c_bstride_zero) { + // CAUSAL-CLOAD (concept receipt: solver i34-9, submission d0ccbe3c). + // A row stride of N + 1 on a bf16 addmm source operand cannot arise + // from any contiguous or broadcast operand of the declared output + // width; it is the deliberate signature of the composed-prefill + // causal bias view, and of nothing else. Synthesize that operand's + // two constants instead of loading them: widened bfloat16 lowest + // finite (0xFF7F) strictly above the causal diagonal placed at + // N - M, widened bfloat16 negative zero on and below it. The addend + // enters through the same TransformAdd, per accumulator element, with + // the same widening as the loaded operand it replaces, so every + // stored word is bit-identical. The padded backing store keeps every + // non-synthesizing branch load-correct at this stride. + const int diag = params->N - params->M; + const int row0 = c_row + mma_op.sm; + const int col0 = c_col + mma_op.sn; + const AccumType mask_add = + static_cast(as_type(0xFF7F0000u)); + const AccumType pass_add = static_cast(-0.0f); + STEEL_PRAGMA_UNROLL + for (short i = 0; i < mma_t::TM; i++) { + STEEL_PRAGMA_UNROLL + for (short j = 0; j < mma_t::TN; j++) { + thread auto& accum = mma_op.Ctile.frag_at(i, j); + const int row = row0 + i * mma_t::TM_stride; + const int col = col0 + j * mma_t::TN_stride; + STEEL_PRAGMA_UNROLL + for (short k = 0; k < decltype(mma_op.Ctile)::kElemsPerFrag; k++) { + accum[k] = epilogue_op_add.apply( + accum[k], (col + k) - row <= diag ? pass_add : mask_add); + } + } + } + synthesized = true; + } + } + if (!synthesized) { + mma_op.apply_epilogue( + C, addmm_params->ldc, addmm_params->fdc, epilogue_op_add); + } } } diff --git a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h index 8267707ade..a00a4e37c2 100644 --- a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h +++ b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h @@ -72,6 +72,49 @@ void gemm_epilogue( }); } +// CAUSAL-CLOAD synthesized epilogue: adds the composed-prefill causal-bias +// constants the loaded operand would have supplied, per accumulator element, +// from the output coordinates the tile already knows. Element order, insert +// point, and widening match gemm_epilogue's non-axpby path exactly. +// clang-format off +template +void gemm_epilogue_causal_synth( + thread NAXTile_t& Dtile, + const int row0, + const int col0, + const int diag) { // clang-format on + using V = typename NAXTile_t::elem_type; + + constexpr short TM = NAXTile_t::kTileRows; + constexpr short TN = NAXTile_t::kTileCols; + + using CFrag = typename NAXTile_t::NAXFrag_t; + + const short2 sc = CFrag::get_coord(); + const V mask_add = static_cast(as_type(0xFF7F0000u)); + const V pass_add = static_cast(-0.0f); + + const_for_loop<0, TM, 1>([&](auto mm) { + const_for_loop<0, TN, 1>([&](auto nn) { + thread auto& delems = Dtile.template frag_at(); + + const int mbase = row0 + sc.y + int(mm) * CFrag::kFragRows; + const int nbase = col0 + sc.x + int(nn) * CFrag::kFragCols; + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < CFrag::kElemRows; i++) { + const int row = mbase + i * CFrag::kElemRowsJump; + STEEL_PRAGMA_UNROLL + for (short j = 0; j < CFrag::kElemCols; j++) { + const int col = nbase + j; + delems[i * CFrag::kElemCols + j] += + (col - row <= diag) ? pass_add : mask_add; + } + } + }); + }); +} + // clang-format off template < typename T, @@ -104,6 +147,11 @@ template < return; } + // CAUSAL-CLOAD eligibility (see steel_gemm_fused.h). + constexpr bool kCausalBiasSynthEligible = + !transpose_a && transpose_b && metal::is_same_v; + bool c_bstride_zero = true; + // Adjust for batch if (has_batch) { const constant auto* A_bstrides = batch_strides; @@ -118,6 +166,9 @@ template < if (use_out_source) { const constant auto* C_bstrides = B_bstrides + params->batch_ndim; C += elem_to_loc(tid.z, batch_shape, C_bstrides, params->batch_ndim); + for (int d = 0; d < params->batch_ndim; d++) { + c_bstride_zero = c_bstride_zero && (C_bstrides[d] == 0); + } } } else { A += params->batch_stride_a * tid.z; @@ -125,6 +176,7 @@ template < if (use_out_source) { C += addmm_params->batch_stride_c * tid.z; + c_bstride_zero = addmm_params->batch_stride_c == 0; } } @@ -203,8 +255,24 @@ template < if ((kAlignedM.value || sgp_sm > 0) && (kAlignedN.value || sgp_sn > 0)) { if (use_out_source) { - gemm_epilogue( - Dtile, C, params, addmm_params, sgp_sm, sgp_sn); + bool synthesized = false; + if constexpr (kCausalBiasSynthEligible) { + if (!do_axpby && kAlignedM.value && kAlignedN.value && + addmm_params->fdc == 1 && + addmm_params->ldc == params->N + 1 && + params->M <= params->N && c_bstride_zero) { + // CAUSAL-CLOAD: signature and exactness argument in + // steel_gemm_fused.h; the synthesized addend is bit-identical + // to the loaded one on every element. + gemm_epilogue_causal_synth( + Dtile, c_row + tm, c_col + tn, params->N - params->M); + synthesized = true; + } + } + if (!synthesized) { + gemm_epilogue( + Dtile, C, params, addmm_params, sgp_sm, sgp_sn); + } } if constexpr (kAlignedM && kAlignedN) { Dtile.store(D, int(params->ldd)); From 201aa99a2eb0fc614a780e455e1f9114b30e4bc4 Mon Sep 17 00:00:00 2001 From: David Tai <8346495+davidtai@users.noreply.github.com> Date: Thu, 3 Sep 2026 10:32:28 -0500 Subject: [PATCH 03/11] perf(metal): fragment-row segment elision for the NAX gather-QMM RHS kernel MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 <11647194+Amal-David@users.noreply.github.com> Co-authored-by: i34-9 <313589706+i34-9@users.noreply.github.com> Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01NS86aXV85ai4eLRw1ADEyX --- mlx/backend/metal/kernels/quantized_nax.h | 324 +++++++++++++++++++--- 1 file changed, 286 insertions(+), 38 deletions(-) diff --git a/mlx/backend/metal/kernels/quantized_nax.h b/mlx/backend/metal/kernels/quantized_nax.h index ed32eb59a7..4a31c7e800 100644 --- a/mlx/backend/metal/kernels/quantized_nax.h +++ b/mlx/backend/metal/kernels/quantized_nax.h @@ -932,6 +932,38 @@ METAL_FUNC void adjust_matrix_offsets( y += tid.z * output_stride; } +// DARKBLOOM GEMMA4 NAX QMM-T ROW-STRIP TILING. +// qmm_t_nax_tgp_impl covers a BM x BN output tile with WM x WN simdgroups. +// The launch shape is fixed by the host (32, WN, WM) and the host is not +// editable, so the threadgroup is always 4 simdgroups over a 64 x 64 tile. +// Stock splits that tile 2 x 2, so each simdgroup owns 32 rows x 32 cols and +// the two simdgroups that share a row band each fetch the SAME 32 rows of the +// activation operand from device memory: A is read twice per threadgroup per +// K step. This constant instead lays the same 4 simdgroups out as 4 row +// strips of 16 rows x 64 cols. The strips are disjoint in M, so every A +// fragment is fetched exactly once, and the B operand -- which already lives +// in threadgroup memory as Ws -- is read wider instead. +// +// Nothing about the K loop moves. BK, SK and TK are untouched, the k and kk1 +// loops keep their bounds and their order, and every output element still +// accumulates over exactly the same k values in exactly the same sequence. +// Only which simdgroup owns an element, and how the owner's fragments are +// shaped, change. The MMA op count per threadgroup is invariant as well: +// stock issues WM*WN * (TM * TN/2 * TK) = 4 * (2 * 1 * 2) = 16 ops per kk1 +// step, the strip layout issues 4 * (1 * 2 * 2) = 16. Both shapes enter the +// same TN-even branch of tile_matmad_nax, so the per-element fragment +// accumulation chain is instruction-for-instruction the same. +// +// The kernel's template parameters, and therefore every kernel-name string +// the host builds, are untouched: BM, BN, BK, WM and WN all keep their +// values and only the interior mapping is re-derived from them. +// +// Kill switch: build with -DDARKBLOOM_GEMMA4_NAX_TILING=0 and SGM/SGN fold +// back to WM/WN, which reproduces the shipped expressions byte for byte. +#ifndef DARKBLOOM_GEMMA4_NAX_TILING +#define DARKBLOOM_GEMMA4_NAX_TILING 1 +#endif + template < typename T, const int group_size, @@ -993,16 +1025,36 @@ METAL_FUNC void qmm_t_nax_tgp_impl( // Make the weight loader loader_w_t loader_w(wl, scales, biases, K, Ws, simd_gid, simd_lid); - constexpr short SM = BM / WM; - constexpr short SN = BN / WN; + // Simdgroup grid over the BM x BN tile. Stock is WM x WN; the row-strip + // layout stacks the same WM*WN simdgroups in the row direction only, so + // no two of them share a row band. See the note on the enable above. + // A row strip is one 16 row NAX fragment row per simdgroup, so the layout + // needs BM >= WM * WN * 16. MLX 0.32.2 also instantiates this kernel with + // BM = 32 ("use smaller bm for many experts and few tokens"), where four + // strips of 16 rows do not fit in the tile; that shape keeps the stock + // WM x WN split, which is what the tile has always used. + constexpr bool kRowStrip = + (DARKBLOOM_GEMMA4_NAX_TILING != 0) && (BM >= WM * WN * 16); + constexpr int SGM = kRowStrip ? (WM * WN) : WM; + constexpr int SGN = kRowStrip ? 1 : WN; + static_assert(SGM * SGN == WM * WN, "simdgroup count must be preserved"); + static_assert(BM % (SGM * 16) == 0, "row strip must be a fragment multiple"); + static_assert(BN % (SGN * 16) == 0, "col strip must be a fragment multiple"); + + constexpr short SM = BM / SGM; + constexpr short SN = BN / SGN; constexpr short SK = 32; constexpr short TM = SM / 16; constexpr short TN = SN / 16; constexpr short TK = SK / 16; - const short tm = SM * (simd_gid / WN); - const short tn = SN * (simd_gid % WN); + // tile_matmad_nax has no branch for an odd TN greater than one; it would + // silently emit no MMA at all. Refuse to compile such a layout. + static_assert(TN == 1 || TN % 2 == 0, "TN must be 1 or even for NAX MMA"); + + const short tm = SM * (simd_gid / SGN); + const short tn = SN * (simd_gid % SGN); constexpr bool transpose_a = false; constexpr bool transpose_b = true; @@ -1462,6 +1514,114 @@ template < w, scales, biases, x, y, Ws, K, N, M, tid, lid, simd_gid, simd_lid); } +// Expert-segment elision for affine_gather_qmm_rhs_nax: the per-tile +// segment loop re-runs the full K-loop once per distinct expert in the +// row tile and discards out-of-segment rows at store_slice. The helpers +// below let a simdgroup skip A loads and MMA for 16-row NAX fragment +// rows that fall wholly outside the current segment's stored row band. +// Fragment rows are independent accumulators, so eliding rows that are +// never stored cannot change any stored element's accumulation sequence. +// Compile-time source constant by design: an enable must never ride a +// function constant magnitude (pipeline-key law). +MLX_MTL_CONST bool kGatherRhsSegmentElide = true; +MLX_MTL_CONST bool kGatherRhsSortedEndpointElide = true; + +// Loads one 16-row fragment row of an A tile from device memory. The +// address arithmetic matches NAXTile::load exactly for that fragment row +// (row offset mm * kFragRows), so the loaded values are identical to the +// full-tile load for the surviving rows. +template +METAL_FUNC void gather_rhs_load_frag_row( + const short mm, + thread ATile& Atile, + const device U* src, + const int ld) { + STEEL_PRAGMA_UNROLL + for (short kk = 0; kk < ATile::kTileCols; ++kk) { + ATile::NAXFrag_t::load( + Atile.frag_at(mm, kk), + src, + ld, + Int<1>{}, + short(mm * ATile::kFragRows), + short(kk * ATile::kFragCols)); + } +} + +// Issues the mm-th fragment row's MMA op sequence of tile_matmad_nax's +// TN-even branch, unchanged: same operands, same per-fragment +// accumulation chain, only the dead fragment rows' ops are absent. +template +METAL_FUNC void gather_rhs_mma_frag_row( + const short mm, + thread CTile& C, + thread ATile& A, + thread BTile& B, + metal::bool_constant tb) { + constexpr short TN = CTile::kTileCols; + constexpr short TK = transpose_b ? BTile::kTileCols : BTile::kTileRows; + constexpr auto ta = metal::bool_constant{}; + static_assert(TN % 2 == 0, "Segment elision expects even TN"); + STEEL_PRAGMA_UNROLL + for (short nn = 0; nn < TN; nn += 2) { + STEEL_PRAGMA_UNROLL + for (short kk = 0; kk < TK; ++kk) { + CTile::NAXFrag_t::mma( + C.frag_at(mm, nn), + C.frag_at(mm, nn + 1), + A.frag_at(mm, kk, ta), + ta, + B.frag_at(kk, nn, tb), + B.frag_at(kk, nn + 1, tb), + tb); + } + } +} + +// DARKBLOOM GEMMA4 NAX GATHER-RHS ROW-STRIP TILING. +// affine_gather_qmm_rhs_nax covers a BM x BN output tile with WM x WN +// simdgroups. The launch shape is fixed by the host (32, WN, WM) and the host +// is not editable, so the threadgroup is always 4 simdgroups over a 64 x 64 +// tile. Stock splits that tile 2 x 2, so each simdgroup owns 32 rows x 32 +// cols and the two simdgroups that share a row band each fetch the SAME 32 +// rows of the activation operand from device memory: A is read twice per +// threadgroup per K step. This constant instead lays the same 4 simdgroups +// out as 4 row strips of 16 rows x 64 cols. The strips are disjoint in M, so +// every A fragment is fetched exactly once, and the B operand -- which +// already lives in threadgroup memory as Ws -- is read wider instead. +// +// Nothing about the K loop moves. BK, SK and TK are untouched, the k, kk1 and +// k_remain loops keep their bounds and their order, and every output element +// still accumulates over exactly the same k values in exactly the same +// sequence. Only which simdgroup owns an element, and how the owner's +// fragments are shaped, change. +// +// COMPOSITION WITH THE SEGMENT ELISION ON THIS KERNEL. The elision is +// expressed at Dtile.kFragRows (16 row) granularity and stays at exactly that +// granularity here: stock gives a simdgroup TM = 2 fragment rows of a 32 row +// band, the strip layout gives TM = 1 fragment row of a 16 row band, and the +// union over the 4 simdgroups is the same 64 rows either way. The live-band +// guard fr < seg_hi && fr + kFragRows > seg_lo tests fr and seg_lo/seg_hi in +// the same tm-relative frame in both layouts, so it decides the same +// intersection of absolute rows against the same segment. offset and +// offset_next stay threadgroup uniform, seg_lo/seg_hi stay simdgroup uniform, +// and gather_rhs_mma_frag_row keeps issuing exactly the TN-even op sequence +// of the shared helper, so the partial-band path and the full path still +// agree op for op. Narrowing the band from 32 rows to 16 can only move a band +// from partial to whole or to empty; it can never make a whole band partial, +// so the elision's own correctness argument is unweakened. +// +// The kernel's template parameters, and therefore every kernel-name string +// the host builds, are untouched: BM, BN, BK, WM and WN all keep their values +// and only the interior mapping is re-derived from them. +// +// Kill switch: build with -DDARKBLOOM_GEMMA4_NAX_GATHER_TILING=0 and SGM/SGN +// fold back to WM/WN, reproducing the shipped expressions byte for byte. +// Independent of the qmm-t family's switch. +#ifndef DARKBLOOM_GEMMA4_NAX_GATHER_TILING +#define DARKBLOOM_GEMMA4_NAX_GATHER_TILING 1 +#endif + template < typename T, int group_size, @@ -1532,16 +1692,39 @@ template < scales += transpose ? y_col_long * K_g : y_col / group_size; biases += transpose ? y_col_long * K_g : y_col / group_size; - constexpr short SM = BM / WM; - constexpr short SN = BN / WN; + // Simdgroup grid over the BM x BN tile. Stock is WM x WN; the row-strip + // layout stacks the same WM*WN simdgroups in the row direction only, so no + // two of them share a row band. See the note on the enable above, including + // why this leaves the segment elision's granularity and guard unchanged. + // A row strip is one 16 row NAX fragment row per simdgroup, so the layout + // needs BM >= WM * WN * 16. MLX 0.32.2 also instantiates this kernel with + // BM = 32 ("use smaller bm for many experts and few tokens"), where four + // strips of 16 rows do not fit in the tile; that shape keeps the stock + // WM x WN split. The segment elision below is unaffected either way: it is + // expressed at Dtile.kFragRows granularity against tm-relative seg_lo / + // seg_hi, which both layouts derive from the same SM. + constexpr bool kRowStrip = + (DARKBLOOM_GEMMA4_NAX_GATHER_TILING != 0) && (BM >= WM * WN * 16); + constexpr int SGM = kRowStrip ? (WM * WN) : WM; + constexpr int SGN = kRowStrip ? 1 : WN; + static_assert(SGM * SGN == WM * WN, "simdgroup count must be preserved"); + static_assert(BM % (SGM * 16) == 0, "row strip must be a fragment multiple"); + static_assert(BN % (SGN * 16) == 0, "col strip must be a fragment multiple"); + + constexpr short SM = BM / SGM; + constexpr short SN = BN / SGN; constexpr short SK = 32; constexpr short TM = SM / 16; constexpr short TN = SN / 16; constexpr short TK = SK / 16; - const short tm = SM * (simd_group_id / WN); - const short tn = SN * (simd_group_id % WN); + // gather_rhs_mma_frag_row issues the shared helper's TN-even op sequence and + // has no branch for an odd TN; an odd TN would silently emit no arithmetic. + static_assert(TN % 2 == 0, "gather segment elision requires an even TN"); + + const short tm = SM * (simd_group_id / SGN); + const short tn = SN * (simd_group_id % SGN); const short sgp_sm = align_M ? SM : min(SM, short(max(0, (M - (y_row + tm))))); @@ -1567,24 +1750,44 @@ template < offset = offset_next; index = index_next; offset_next = tgp_bm; - for (; n < tgp_bm; n++) { - if (indices[y_row + n] != index) { - offset_next = n; - index_next = indices[y_row + n]; - break; + // gather_qmm_rhs is dispatched only for right-sorted indices. If this + // segment's expert matches the tile endpoint, sortedness proves that the + // remaining suffix is one segment and the per-row probe can stop here. + if (kGatherRhsSortedEndpointElide && + indices[y_row + tgp_bm - 1] == index) { + n = tgp_bm; + } else { + for (; n < tgp_bm; n++) { + if (indices[y_row + n] != index) { + offset_next = n; + index_next = indices[y_row + n]; + break; + } } } threadgroup_barrier(mem_flags::mem_none); - const short m_lo_lim = min(int(sgp_sm), max(0, offset - tm)); - const short m_hi_lim = min(int(sgp_sm), max(0, offset_next - tm)); - const bool sg_active = m_hi_lim > m_lo_lim; - NAXTile Dtile; Dtile.clear(); const device T* xn = x + tm * K; + // This simdgroup's stored row band for the current expert segment, + // hoisted ahead of the K-loop (it depends only on offset, offset_next, + // tm and sgp_sm, all known here). The stock path computes the full + // tile and discards rows outside [seg_lo, seg_hi) at store_slice; with + // the elision enabled those rows' A loads and MMA ops are skipped + // instead. Cooperative weight loads and every threadgroup_barrier stay + // unconditional, so barrier convergence is preserved, and seg_* are + // uniform within a simdgroup (offset/offset_next are threadgroup + // uniform). With the enable off both flags fold to false and only the + // stock path below runs. + const short seg_lo = min(int(sgp_sm), max(0, offset - tm)); + const short seg_hi = min(int(sgp_sm), max(0, offset_next - tm)); + const bool seg_empty = kGatherRhsSegmentElide && (seg_hi <= seg_lo); + const bool seg_partial = kGatherRhsSegmentElide && !seg_empty && + !(seg_lo == 0 && seg_hi == sgp_sm); + // Prepare threadgroup loading operations thread loader_w_t loader_w( wl + index * stride_w, @@ -1608,9 +1811,44 @@ template < threadgroup_barrier(mem_flags::mem_threadgroup); - STEEL_PRAGMA_NO_UNROLL - for (int kk1 = 0; kk1 < BK; kk1 += SK) { - if (sg_active) { + if (seg_partial && kAlignedM.value) { + // 16-row fragment-row granularity: only fragment rows that + // intersect [seg_lo, seg_hi) load A and issue MMA. Each live + // fragment row runs the exact op sequence of the stock path. + STEEL_PRAGMA_NO_UNROLL + for (int kk1 = 0; kk1 < BK; kk1 += SK) { + NAXTile Atile; + NAXTile Btile; + + volatile int compiler_barrier; + + if constexpr (transpose) { + Btile.template load( + Ws + tn * BK_padded + kk1); + } else { + Btile.template load( + Ws + tn + kk1 * BN_padded); + } + + STEEL_PRAGMA_UNROLL + for (short mm = 0; mm < TM; mm++) { + const short fr = short(mm * Dtile.kFragRows); + if (fr < seg_hi && short(fr + Dtile.kFragRows) > seg_lo) { + gather_rhs_load_frag_row(mm, Atile, xn + kk1, K); + gather_rhs_mma_frag_row( + mm, + Dtile, + Atile, + Btile, + metal::bool_constant{}); + } + } + + (void)compiler_barrier; + } + } else if (!seg_empty) { + STEEL_PRAGMA_NO_UNROLL + for (int kk1 = 0; kk1 < BK; kk1 += SK) { NAXTile Atile; NAXTile Btile; @@ -1623,9 +1861,11 @@ template < } if constexpr (transpose) { - Btile.template load(Ws + tn * BK_padded + kk1); + Btile.template load( + Ws + tn * BK_padded + kk1); } else { - Btile.template load(Ws + tn + kk1 * BN_padded); + Btile.template load( + Ws + tn + kk1 * BN_padded); } tile_matmad_nax( @@ -1648,9 +1888,12 @@ template < loader_w.load_safe(tile_w); threadgroup_barrier(mem_flags::mem_threadgroup); - STEEL_PRAGMA_NO_UNROLL - for (int kk1 = 0; kk1 < BK; kk1 += SK) { - if (sg_active) { + // Elision here is band-granular only (seg_empty): a partial band + // runs the stock tail, whose extra MMA lands in fragment rows + // that are never stored. + if (!seg_empty) { + STEEL_PRAGMA_NO_UNROLL + for (int kk1 = 0; kk1 < BK; kk1 += SK) { NAXTile Atile; NAXTile Btile; @@ -1660,9 +1903,11 @@ template < Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); if constexpr (transpose) { - Btile.template load(Ws + tn * BK_padded + kk1); + Btile.template load( + Ws + tn * BK_padded + kk1); } else { - Btile.template load(Ws + tn + kk1 * BN_padded); + Btile.template load( + Ws + tn + kk1 * BN_padded); } tile_matmad_nax( @@ -1679,20 +1924,23 @@ template < threadgroup_barrier(mem_flags::mem_threadgroup); - // Store results to device memory - if constexpr (kAlignedN.value) { - if (m_lo_lim == 0 && m_hi_lim == SM) { - Dtile.store(y + tm * N + tn, N); + // Store results to device memory. seg_lo/seg_hi are the stock + // m_lo_lim/m_hi_lim, hoisted ahead of the K-loop. + if (!seg_empty) { + if constexpr (kAlignedN.value) { + if (seg_lo == 0 && seg_hi == SM) { + Dtile.store(y + tm * N + tn, N); + } else { + Dtile.store_slice( + y + tm * N + tn, N, short2(0, seg_lo), short2(SN, seg_hi)); + } } else { Dtile.store_slice( - y + tm * N + tn, N, short2(0, m_lo_lim), short2(SN, m_hi_lim)); + y + tm * N + tn, + N, + short2(0, seg_lo), + short2(sgp_sn, seg_hi)); } - } else { - Dtile.store_slice( - y + tm * N + tn, - N, - short2(0, m_lo_lim), - short2(sgp_sn, m_hi_lim)); } }); }); From f76fb7e46f550b7a30efcc0d9be17a04659ba43b Mon Sep 17 00:00:00 2001 From: David Tai <8346495+davidtai@users.noreply.github.com> Date: Thu, 3 Sep 2026 10:34:31 -0500 Subject: [PATCH 04/11] perf(metal): the Gemma 4 affine-quantized QMV tier family in quantized.h MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 `. 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 <95952950+0xkydo@users.noreply.github.com> Co-authored-by: Amal-David <11647194+Amal-David@users.noreply.github.com> Co-authored-by: DashiellB <65423051+DashiellB@users.noreply.github.com> Co-authored-by: brandonegg <13079136+brandonegg@users.noreply.github.com> Co-authored-by: delordemm1 <46292455+delordemm1@users.noreply.github.com> Co-authored-by: ercumentyildirim <43972346+ercumentyildirim@users.noreply.github.com> Co-authored-by: exakoss <67432899+exakoss@users.noreply.github.com> Co-authored-by: i34-9 <313589706+i34-9@users.noreply.github.com> Co-authored-by: ivanfioravanti <1069210+ivanfioravanti@users.noreply.github.com> Co-authored-by: jungjipdo <130676635+jungjipdo@users.noreply.github.com> Co-authored-by: newjordan <11369410+newjordan@users.noreply.github.com> Co-authored-by: polymorf <127736+polymorf@users.noreply.github.com> Co-authored-by: rinaldofesta <5622471+rinaldofesta@users.noreply.github.com> Co-authored-by: rube-de <8930910+rube-de@users.noreply.github.com> Co-authored-by: samfenwick <45273188+samfenwick@users.noreply.github.com> Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01NS86aXV85ai4eLRw1ADEyX --- mlx/backend/metal/kernels/quantized.h | 2336 ++++++++++++++++++++++++- 1 file changed, 2281 insertions(+), 55 deletions(-) diff --git a/mlx/backend/metal/kernels/quantized.h b/mlx/backend/metal/kernels/quantized.h index 6831cfd294..08f73eb8b1 100644 --- a/mlx/backend/metal/kernels/quantized.h +++ b/mlx/backend/metal/kernels/quantized.h @@ -289,6 +289,197 @@ inline U qdot( return scale * accum + sum * bias; } +// One affine-4 dot product against a packed weight vector ALREADY held in +// registers. Byte-for-byte the bits == 4 arm of qdot: same nibble masks, same +// four-term expression, same accumulation order over i, same +// `scale * accum + sum * bias` close. Only the residence of `w` differs +// (thread instead of device), which is what lets one weight fetch serve four +// cohort input rows without holding four rows of x live at once. +template +inline U qdot_affine4_registered( + const thread uint16_t* w, + const thread U* x_thread, + U scale, + U bias, + U sum) { + U accum = 0; + for (int i = 0; i < (values_per_thread / 4); i++) { + accum += + (x_thread[4 * i] * (w[i] & 0x000f) + + x_thread[4 * i + 1] * (w[i] & 0x00f0) + + x_thread[4 * i + 2] * (w[i] & 0x0f00) + + x_thread[4 * i + 3] * (w[i] & 0xf000)); + } + return scale * accum + sum * bias; +} + +// Consume the same two adjacent packed uint16 values through one aligned +// 32-bit device load. The low and high halves retain the original arithmetic +// order while halving the explicit weight-load instructions. +template +inline U qdot_affine4_registered_word( + uint packed_word, + const thread U* x_thread, + U scale, + U bias, + U sum) { + static_assert(values_per_thread == 8, "Word load expects eight 4-bit values"); + const uint packed0 = packed_word & 0xffffu; + const uint packed1 = packed_word >> 16; + U accum = + (x_thread[0] * (packed0 & 0x000f) + + x_thread[1] * (packed0 & 0x00f0) + + x_thread[2] * (packed0 & 0x0f00) + + x_thread[3] * (packed0 & 0xf000)); + accum += + (x_thread[4] * (packed1 & 0x000f) + + x_thread[5] * (packed1 & 0x00f0) + + x_thread[6] * (packed1 & 0x0f00) + + x_thread[7] * (packed1 & 0xf000)); + return scale * accum + sum * bias; +} + +// Two independent affine-4 dot products over one packed weight vector. Each +// accumulator retains the scalar qdot operation order; only the packed weight +// load is shared between adjacent assignments routed to the same expert. +template +inline void qdot_affine4_pair( + const device uint8_t* w, + const thread U* x0, + const thread U* x1, + U scale, + U bias, + U sum0, + U sum1, + thread U& out0, + thread U& out1) { + static_assert(values_per_thread == 8, "Word load expects eight 4-bit values"); + const uint packedWord = *((const device uint*)w); + const uint packed0 = packedWord & 0xffffu; + const uint packed1 = packedWord >> 16; + U accum0 = + (x0[0] * (packed0 & 0x000f) + + x0[1] * (packed0 & 0x00f0) + + x0[2] * (packed0 & 0x0f00) + + x0[3] * (packed0 & 0xf000)); + U accum1 = + (x1[0] * (packed0 & 0x000f) + + x1[1] * (packed0 & 0x00f0) + + x1[2] * (packed0 & 0x0f00) + + x1[3] * (packed0 & 0xf000)); + accum0 += + (x0[4] * (packed1 & 0x000f) + + x0[5] * (packed1 & 0x00f0) + + x0[6] * (packed1 & 0x0f00) + + x0[7] * (packed1 & 0xf000)); + accum1 += + (x1[4] * (packed1 & 0x000f) + + x1[5] * (packed1 & 0x00f0) + + x1[6] * (packed1 & 0x0f00) + + x1[7] * (packed1 & 0xf000)); + out0 = scale * accum0 + sum0 * bias; + out1 = scale * accum1 + sum1 * bias; +} + +// Two independent affine-4 dot products over one register-held packed 32-bit word. +template +inline void qdot_affine4_pair_word( + uint packedWord, + const thread U* x0, + const thread U* x1, + U scale, + U bias, + U sum0, + U sum1, + thread U& out0, + thread U& out1) { + static_assert(values_per_thread == 8, "Word load expects eight 4-bit values"); + const uint packed0 = packedWord & 0xffffu; + const uint packed1 = packedWord >> 16; + U accum0 = + (x0[0] * (packed0 & 0x000f) + + x0[1] * (packed0 & 0x00f0) + + x0[2] * (packed0 & 0x0f00) + + x0[3] * (packed0 & 0xf000)); + U accum1 = + (x1[0] * (packed0 & 0x000f) + + x1[1] * (packed0 & 0x00f0) + + x1[2] * (packed0 & 0x0f00) + + x1[3] * (packed0 & 0xf000)); + accum0 += + (x0[4] * (packed1 & 0x000f) + + x0[5] * (packed1 & 0x00f0) + + x0[6] * (packed1 & 0x0f00) + + x0[7] * (packed1 & 0xf000)); + accum1 += + (x1[4] * (packed1 & 0x000f) + + x1[5] * (packed1 & 0x00f0) + + x1[6] * (packed1 & 0x0f00) + + x1[7] * (packed1 & 0xf000)); + out0 = scale * accum0 + sum0 * bias; + out1 = scale * accum1 + sum1 * bias; +} + +// One affine-8 dot product against a byte weight vector ALREADY held in +// registers. Byte-for-byte the bits == 8 arm of qdot: same per-element +// multiply, same accumulation order over i, same `scale * accum + sum * bias` +// close. Only the residence of `w` differs (thread instead of device). +template +inline U qdot_affine8_registered( + const thread uint8_t* w, + const thread U* x_thread, + U scale, + U bias, + U sum) { + U accum = 0; + for (int i = 0; i < values_per_thread; i++) { + accum += x_thread[i] * w[i]; + } + return scale * accum + sum * bias; +} + +// The same four products, accumulated in the same order into an accumulator +// opened at zero, over the same four bytes taken from one packed word. The +// byte at the lowest address is the low byte of the word. +template +inline U qdot_affine8_registered_word( + uint packed_word, + const thread U* x_thread, + U scale, + U bias, + U sum) { + U accum = 0; + accum += x_thread[0] * U(packed_word & 0xffu); + accum += x_thread[1] * U((packed_word >> 8) & 0xffu); + accum += x_thread[2] * U((packed_word >> 16) & 0xffu); + accum += x_thread[3] * U(packed_word >> 24); + return scale * accum + sum * bias; +} + +// Two independent affine-8 dot products over one byte weight vector. Keep the +// per-row scalar accumulation order of qdot while sharing each weight load. +template +inline void qdot_affine8_pair( + const device uint8_t* w, + const thread U* x0, + const thread U* x1, + U scale, + U bias, + U sum0, + U sum1, + thread U& out0, + thread U& out1) { + U accum0 = 0; + U accum1 = 0; + for (int i = 0; i < values_per_thread; i++) { + const uint8_t packed = w[i]; + accum0 += x0[i] * packed; + accum1 += x1[i] * packed; + } + out0 = scale * accum0 + sum0 * bias; + out1 = scale * accum1 + sum1 * bias; +} + template inline U qdot_safe( const device uint8_t* w, @@ -821,6 +1012,378 @@ METAL_FUNC void qmv_fast_impl( } } +// Exact-order affine4/g64 multi-row QMV. The frozen host launches M x-groups +// for each 8-output tile. Pair adjacent input rows in one group while keeping +// the stock two-simdgroup by four-output-row layout. Each active group caches a +// weight tile once and applies the stock arithmetic independently to one or two +// inputs; unused host groups return without reading weights. load_vector, the +// qdot expression, K accumulation order, and simd_sum remain identical to +// qmv_fast_impl for every output element. +template +inline U qdot_affine4_loaded( + const thread uint16_t* ws, + const thread U* x_thread, + U scale, + U bias, + U sum) { + U accum = 0; + for (int i = 0; i < 4; i++) { + accum += + (x_thread[4 * i] * (ws[i] & 0x000f) + + x_thread[4 * i + 1] * (ws[i] & 0x00f0) + + x_thread[4 * i + 2] * (ws[i] & 0x0f00) + + x_thread[4 * i + 3] * (ws[i] & 0xf000)); + } + return scale * accum + sum * bias; +} + +inline float2 qdot_affine4_loaded_pair( + const thread uint16_t* ws, + const thread float* x0, + const thread float* x1, + float scale, + float bias, + float2 sum) { + float2 accum = 0; + for (int i = 0; i < 4; i++) { + accum += + (float2(x0[4 * i], x1[4 * i]) * (ws[i] & 0x000f) + + float2(x0[4 * i + 1], x1[4 * i + 1]) * (ws[i] & 0x00f0) + + float2(x0[4 * i + 2], x1[4 * i + 2]) * (ws[i] & 0x0f00) + + float2(x0[4 * i + 3], x1[4 * i + 3]) * (ws[i] & 0xf000)); + } + return scale * accum + sum * bias; +} + +template +METAL_FUNC void qmv_fast_crossrow_affine4_g64( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x, + device T* y, + const constant int& in_vec_size, + const constant int& out_vec_size, + uint3 tid, + uint simd_gid, + uint simd_lid) { + static_assert(M >= 2 && M <= 9, "multi-row QMV supports M in [2, 9]"); + constexpr int inputs_per_group = 2; + constexpr int rows_per_simd = 4; + constexpr int values_per_thread = 16; + constexpr int block_size = values_per_thread * SIMD_SIZE; + constexpr int in_vec_bytes_per_row_divisor = 2; + constexpr int bytes_per_lane = 8; + + const int first_m = int(tid.x) * inputs_per_group; + if (first_m >= M) { + return; + } + const int out_row = int(tid.y) * 8 + int(simd_gid) * rows_per_simd; + const int in_vec_size_w = in_vec_size / in_vec_bytes_per_row_divisor; + const int in_vec_size_g = in_vec_size / 64; + + const bool has_pair = first_m + 1 < M; + thread float2 pair_result[rows_per_simd]; + thread float single_result[rows_per_simd]; + for (int r = 0; r < rows_per_simd; r++) { + pair_result[r] = 0.0f; + single_result[r] = 0.0f; + } + + for (int k = 0; k < in_vec_size; k += block_size) { + thread uint16_t packed[rows_per_simd][4]; + thread float scale_local[rows_per_simd]; + thread float bias_local[rows_per_simd]; + + for (int r = 0; r < rows_per_simd; r++) { + const int row = out_row + r; + const device uint8_t* wb = + reinterpret_cast(w) + + row * in_vec_size_w + k / 2 + simd_lid * bytes_per_lane; + const device uint16_t* ws = + reinterpret_cast(wb); + for (int i = 0; i < 4; i++) { + packed[r][i] = ws[i]; + } + const int group_index = + row * in_vec_size_g + k / 64 + simd_lid / 4; + scale_local[r] = scales[group_index]; + bias_local[r] = biases[group_index]; + } + + thread float x0[values_per_thread]; + const device T* xm0 = + x + first_m * in_vec_size + k + simd_lid * values_per_thread; + const float sum0 = + load_vector(xm0, x0); + if (has_pair) { + thread float x1[values_per_thread]; + const device T* xm1 = xm0 + in_vec_size; + const float sum1 = + load_vector(xm1, x1); + for (int r = 0; r < rows_per_simd; r++) { + pair_result[r] += qdot_affine4_loaded_pair( + packed[r], x0, x1, scale_local[r], bias_local[r], + float2(sum0, sum1)); + } + } else { + for (int r = 0; r < rows_per_simd; r++) { + single_result[r] += qdot_affine4_loaded( + packed[r], x0, scale_local[r], bias_local[r], sum0); + } + } + } + + if (has_pair) { + for (int r = 0; r < rows_per_simd; r++) { + const float reduced0 = simd_sum(pair_result[r].x); + const float reduced1 = simd_sum(pair_result[r].y); + if (simd_lid == 0) { + y[first_m * out_vec_size + out_row + r] = static_cast(reduced0); + y[(first_m + 1) * out_vec_size + out_row + r] = + static_cast(reduced1); + } + } + } else { + for (int r = 0; r < rows_per_simd; r++) { + const float reduced = simd_sum(single_result[r]); + if (simd_lid == 0) { + y[first_m * out_vec_size + out_row + r] = static_cast(reduced); + } + } + } +} + +// Wider row sharing for the affine4/g64 multi-row QMV. Same contract as +// qmv_fast_crossrow_affine4_g64: the frozen host launches M x-groups for each +// 8-output tile, so a group that claims NA adjacent input rows lets the +// remaining host groups return without reading weights. NA up to 4 shares one +// nibble mask and one integer-to-float conversion across NA inputs while +// holding only four x values per input live at a time, so the register +// footprint stays near the two-input kernel's. load_vector, the qdot +// expression, the K accumulation order and simd_sum are unchanged for every +// output element. +template +METAL_FUNC void qmv_fast_crossrow_affine4_g64_wide( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x, + device T* y, + const int in_vec_size, + const int out_vec_size, + int first_m, + int out_row, + uint simd_lid) { + static_assert(NA >= 2 && NA <= 4, "wide multi-row QMV supports NA in [2, 4]"); + typedef vec VF; + constexpr int rows_per_simd = 4; + constexpr int values_per_thread = 16; + constexpr int block_size = values_per_thread * SIMD_SIZE; + constexpr int bytes_per_lane = 8; + const int in_vec_size_w = in_vec_size / 2; + const int in_vec_size_g = in_vec_size / 64; + + VF acc[rows_per_simd]; + for (int r = 0; r < rows_per_simd; r++) { + acc[r] = VF(0.0f); + } + + for (int k = 0; k < in_vec_size; k += block_size) { + thread uint16_t packed[rows_per_simd][4]; + thread float scale_local[rows_per_simd]; + thread float bias_local[rows_per_simd]; + for (int r = 0; r < rows_per_simd; r++) { + const int row = out_row + r; + const device uint16_t* ws = reinterpret_cast( + reinterpret_cast(w) + row * in_vec_size_w + + k / 2 + simd_lid * bytes_per_lane); + for (int i = 0; i < 4; i++) { + packed[r][i] = ws[i]; + } + const int group_index = row * in_vec_size_g + k / 64 + simd_lid / 4; + scale_local[r] = scales[group_index]; + bias_local[r] = biases[group_index]; + } + + VF sums = VF(0.0f); + VF partial[rows_per_simd]; + for (int r = 0; r < rows_per_simd; r++) { + partial[r] = VF(0.0f); + } + for (int i = 0; i < 4; i++) { + VF a0, a1, a2, a3; + for (int m = 0; m < NA; m++) { + const device T* xm = x + (first_m + m) * in_vec_size + k + + simd_lid * values_per_thread + 4 * i; + thread float xc[4]; + if (DIRECT_NIBBLES) { + xc[0] = static_cast(xm[0]); + xc[1] = static_cast(xm[1]); + xc[2] = static_cast(xm[2]); + xc[3] = static_cast(xm[3]); + // Preserve the incumbent BF16 expression tree used for the affine + // bias correction; only the qdot nibble extraction changes. + sums[m] += xm[0] + xm[1] + xm[2] + xm[3]; + } else { + sums[m] += load_vector(xm, xc); + } + a0[m] = xc[0]; + a1[m] = xc[1]; + a2[m] = xc[2]; + a3[m] = xc[3]; + } + for (int r = 0; r < rows_per_simd; r++) { + if (DIRECT_NIBBLES) { + partial[r] += (a0 * (packed[r][i] & 0x000f) + + a1 * ((packed[r][i] >> 4) & 0x000f) + + a2 * ((packed[r][i] >> 8) & 0x000f) + + a3 * ((packed[r][i] >> 12) & 0x000f)); + } else { + partial[r] += (a0 * (packed[r][i] & 0x000f) + + a1 * (packed[r][i] & 0x00f0) + + a2 * (packed[r][i] & 0x0f00) + + a3 * (packed[r][i] & 0xf000)); + } + } + } + for (int r = 0; r < rows_per_simd; r++) { + acc[r] += scale_local[r] * partial[r] + sums * bias_local[r]; + } + } + + for (int r = 0; r < rows_per_simd; r++) { + for (int m = 0; m < NA; m++) { + const float reduced = simd_sum(acc[r][m]); + if (simd_lid == 0) { + y[(first_m + m) * out_vec_size + out_row + r] = + static_cast(reduced); + } + } + } +} + +// Single-row (M == 1) affine2/g64 fast QMV for the coarse compact draft +// readout (out_vec_size == 98_336, bits == 2) of the promoted draft-rerank +// scheme, at 32 values per lane: each lane loads ONE uint64 (32 packed +// 2-bit values) per row per k-block, halving load count and k-blocks versus +// the generic 16-value form. Duo values are extracted by shift and +// multiplied by the UNSCALED activation: (x / 4^k) * (w & (3 << 2k)) and +// x * ((w >> 2k) & 3) are the same real product (power-of-two scaling is +// exact in FP32), so every elementary product equals the generic +// qmv_fast_impl value; the wider lane coverage reassociates the +// FP32 partial sums, which is safe for this stage because the coarse +// shortlist is approximate by design and the exact affine-4 rerank plus +// target verification decide every emitted token. The serial leg runs no +// 2-bit matmul (all its projections are affine-4), and out_vec_size == +// 98_336 exists only in the compact draft readout, so the dispatch gate +// below cannot touch the serial numerator or the denominator band. +template +METAL_FUNC void qmv_fast_singlerow_affine2_g64( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x, + device T* y, + const constant int& in_vec_size, + const constant int& out_vec_size, + uint3 tid, + uint simd_gid, + uint simd_lid) { + constexpr int rows_per_simd = 4; + constexpr int values_per_thread = 32; + constexpr int block_size = values_per_thread * SIMD_SIZE; + constexpr int bytes_per_lane = 8; // 32 values x 2 bits = 8 bytes + const int in_vec_size_w = in_vec_size / 4; // weight bytes per output row + const int in_vec_size_g = in_vec_size / 64; // scale groups per output row + + const int out_row = int(tid.y) * 8 + int(simd_gid) * rows_per_simd; + + thread float result[rows_per_simd]; + for (int r = 0; r < rows_per_simd; r++) { + result[r] = 0.0f; + } + + for (int k = 0; k < in_vec_size; k += block_size) { + thread ulong packed[rows_per_simd]; + thread float scale_local[rows_per_simd]; + thread float bias_local[rows_per_simd]; + for (int r = 0; r < rows_per_simd; r++) { + const int row = out_row + r; + const device uint8_t* ws = reinterpret_cast(w) + + row * in_vec_size_w + k / 4 + simd_lid * bytes_per_lane; + packed[r] = *reinterpret_cast(ws); + // 32 values per lane = half of one 64-value group. + const int group_index = + row * in_vec_size_g + k / 64 + (simd_lid * values_per_thread) / 64; + scale_local[r] = scales[group_index]; + bias_local[r] = biases[group_index]; + } + + thread float x0[values_per_thread]; + const device T* xm = x + k + simd_lid * values_per_thread; + float sum = 0.0f; + for (int i = 0; i < values_per_thread; i += 4) { + x0[i] = static_cast(xm[i]); + x0[i + 1] = static_cast(xm[i + 1]); + x0[i + 2] = static_cast(xm[i + 2]); + x0[i + 3] = static_cast(xm[i + 3]); + sum += xm[i] + xm[i + 1] + xm[i + 2] + xm[i + 3]; + } + + for (int r = 0; r < rows_per_simd; r++) { + float accum = 0.0f; + #pragma unroll + for (int j = 0; j < 32; j++) { + accum += x0[j] * float((packed[r] >> (2 * j)) & 0x03ul); + } + result[r] += scale_local[r] * accum + sum * bias_local[r]; + } + } + + for (int r = 0; r < rows_per_simd; r++) { + const float reduced = simd_sum(result[r]); + if (simd_lid == 0) { + y[out_row + r] = static_cast(reduced); + } + } +} + +// IPG = ceil(M / ceil(M / 4)): the fewest weight streams reachable at NA <= 4, +// with the remainder spread evenly so no group runs a one-row tail. +template +METAL_FUNC void qmv_fast_crossrow_affine4_g64_m( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x, + device T* y, + const constant int& in_vec_size, + const constant int& out_vec_size, + uint3 tid, + uint simd_gid, + uint simd_lid) { + static_assert(M >= 3 && M <= 9, "wide multi-row QMV dispatch covers M in [3, 9]"); + static_assert(M % IPG != 1, "a one-input tail group is not instantiated"); + constexpr int TAIL = M % IPG; + const int first_m = int(tid.x) * IPG; + if (first_m >= M) { + return; + } + const int out_row = int(tid.y) * 8 + int(simd_gid) * 4; + if (TAIL == 0 || M - first_m >= IPG) { + qmv_fast_crossrow_affine4_g64_wide( + w, scales, biases, x, y, in_vec_size, out_vec_size, + first_m, out_row, simd_lid); + } else { + qmv_fast_crossrow_affine4_g64_wide< + T, (TAIL >= 2 ? TAIL : 2), DIRECT_NIBBLES>( + w, scales, biases, x, y, in_vec_size, out_vec_size, + first_m, out_row, simd_lid); + } +} + template METAL_FUNC void qmv_impl( const device uint32_t* w, @@ -872,7 +1435,7 @@ METAL_FUNC void qmv_impl( y += tid.x * out_vec_size + out_row; int k = 0; - for (; k < in_vec_size - block_size; k += block_size) { + for (; k <= in_vec_size - block_size; k += block_size) { U sum = load_vector(x, x_thread); for (int row = 0; @@ -935,7 +1498,7 @@ METAL_FUNC void qmv_impl( y += tid.x * out_vec_size + used_out_row; int k = 0; - for (; k < in_vec_size - block_size; k += block_size) { + for (; k <= in_vec_size - block_size; k += block_size) { U sum = load_vector(x, x_thread); for (int row = 0; row < results_per_simdgroup; row++) { @@ -943,43 +1506,900 @@ METAL_FUNC void qmv_impl( const device T* sl = scales + row * in_vec_size_g; const device T* bl = biases + row * in_vec_size_g; - U s = sl[0]; - U b = bl[0]; - result[row] += - qdot(wl, x_thread, s, b, sum); - } + U s = sl[0]; + U b = bl[0]; + result[row] += + qdot(wl, x_thread, s, b, sum); + } + + ws += block_size * bytes_per_pack / pack_factor; + scales += block_size / group_size; + biases += block_size / group_size; + x += block_size; + } + const int tail_values = static_cast(in_vec_size - k); + if (tail_values > 0) { + // Affine callers keep K a whole number of quantization groups and k + // advances by whole blocks, so the tail is a whole number of + // values_per_thread lane packets: routed-expert down_proj K=704 leaves + // 192 values = 24 complete packets, dense down_proj K=2112 (8-bit) + // leaves 64 = 16. Active lanes run the fixed unrolled loader and qdot; + // the dynamic safe-tail remains only for a genuinely partial packet, + // which no affine caller presents. + if (tail_values % values_per_thread == 0) { + const uint active_tail_lanes = uint(tail_values / values_per_thread); + if (simd_lid < active_tail_lanes) { + U sum = load_vector(x, x_thread); + + for (int row = 0; row < results_per_simdgroup; row++) { + auto wl = (const device uint8_t*)(ws + row * in_vec_size_w); + const device T* sl = scales + row * in_vec_size_g; + const device T* bl = biases + row * in_vec_size_g; + + U s = sl[0]; + U b = bl[0]; + result[row] += + qdot(wl, x_thread, s, b, sum); + } + } + } else { + const int remaining = clamp( + static_cast(tail_values - simd_lid * values_per_thread), + 0, + values_per_thread); + if (remaining > 0) { + U sum = load_vector_safe( + x, x_thread, remaining); + + for (int row = 0; row < results_per_simdgroup; row++) { + auto wl = (const device uint8_t*)(ws + row * in_vec_size_w); + const device T* sl = scales + row * in_vec_size_g; + const device T* bl = biases + row * in_vec_size_g; + + U s = sl[0]; + U b = bl[0]; + result[row] += qdot_safe( + wl, x_thread, s, b, sum, remaining); + } + } + } + } + for (int row = 0; row < results_per_simdgroup; row++) { + result[row] = simd_sum(result[row]); + if (simd_lid == 0) { + y[row] = static_cast(result[row]); + } + } + } +} + +template +METAL_FUNC void qmv_affine4_g64_pair_impl( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x0, + const device T* x1, + device T* y0, + device T* y1, + const constant int& in_vec_size, + uint3 tid [[threadgroup_position_in_grid]], + uint simd_gid [[simdgroup_index_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + constexpr int num_simdgroups = 2; + constexpr int results_per_simdgroup = 4; + constexpr int values_per_thread = 8; + constexpr int block_size = values_per_thread * SIMD_SIZE; + constexpr int bytes_per_thread = 4; + constexpr int scale_step_per_thread = 8; + + const device uint8_t* ws = (const device uint8_t*)w; + thread float x0_thread[values_per_thread]; + thread float x1_thread[values_per_thread]; + thread uint packed[results_per_simdgroup]; + thread float scale_local[results_per_simdgroup]; + thread float bias_local[results_per_simdgroup]; + thread float result0[results_per_simdgroup] = {0}; + thread float result1[results_per_simdgroup] = {0}; + + const int in_vec_size_w = in_vec_size / 2; + const int in_vec_size_g = in_vec_size / 64; + const int out_row = tid.y * (num_simdgroups * results_per_simdgroup) + + simd_gid * results_per_simdgroup; + + ws += out_row * in_vec_size_w + simd_lid * bytes_per_thread; + scales += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + biases += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + x0 += simd_lid * values_per_thread; + x1 += simd_lid * values_per_thread; + y0 += out_row; + y1 += out_row; + + int k = 0; + for (; k <= in_vec_size - block_size; k += block_size) { + for (int row = 0; row < results_per_simdgroup; row++) { + packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); + scale_local[row] = scales[row * in_vec_size_g]; + bias_local[row] = biases[row * in_vec_size_g]; + } + + float sum0 = load_vector(x0, x0_thread); + float sum1 = load_vector(x1, x1_thread); + + for (int row = 0; row < results_per_simdgroup; row++) { + float dot0; + float dot1; + qdot_affine4_pair_word( + packed[row], x0_thread, x1_thread, scale_local[row], bias_local[row], sum0, sum1, dot0, dot1); + result0[row] += dot0; + result1[row] += dot1; + } + + ws += block_size / 2; + scales += block_size / 64; + biases += block_size / 64; + x0 += block_size; + x1 += block_size; + } + + // Every Gemma 4 caller entering this specialized g64 path has K aligned to + // 64. The final block therefore contains an integral number of complete + // eight-value lane packets (32 lanes for K=2816, 24 for expert down_proj + // K=704); no active lane needs the generic dynamic safe-tail loops. + const uint active_tail_lanes = + uint((in_vec_size - k) / values_per_thread); + if (simd_lid < active_tail_lanes) { + for (int row = 0; row < results_per_simdgroup; row++) { + packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); + scale_local[row] = scales[row * in_vec_size_g]; + bias_local[row] = biases[row * in_vec_size_g]; + } + + float sum0 = + load_vector(x0, x0_thread); + float sum1 = + load_vector(x1, x1_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + float dot0; + float dot1; + qdot_affine4_pair_word( + packed[row], x0_thread, x1_thread, scale_local[row], bias_local[row], sum0, sum1, dot0, dot1); + result0[row] += dot0; + result1[row] += dot1; + } + } + + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] = simd_sum(result0[row]); + result1[row] = simd_sum(result1[row]); + if (simd_lid == 0) { + y0[row] = static_cast(result0[row]); + y1[row] = static_cast(result1[row]); + } + } +} + +// Four-row weight-stream sharing for the ordinary (plain-order) affine4/g64 +// QMV, held to a PAIR-SIZED register footprint. Same geometry as +// qmv_affine4_g64_pair_impl: two simdgroups by four output rows per +// 64-thread group, values_per_thread = 8, block_size = 256, +// scale_step_per_thread = 8, same load_vector / load_vector_safe tail. +// +// RESIDENCY is the whole point. The retired four-row quad this replaces held +// all four cohort rows of x live across a block (4 x 8 = 32 floats) on top of +// its 16 accumulators, and measured SLOWER than the two-row pair it was meant +// to beat -- 226.3 us against 208.7 us at N = 8192 on the ranked box, with the +// halved weight stream never converting. This kernel fetches the block's +// packed weights and scale/bias into registers ONCE and then walks the four +// input rows in sequence through a SINGLE eight-value x buffer, so only one +// row of x is live at a time -- the discipline of +// qmv_fast_crossrow_affine4_g64_wide, which keeps four x values per input row +// live and states the same reason. With it the collapse converts: 197.5 us at +// N = 8192, under BOTH the pair kernel and the retired quad. +// +// Exactness is unchanged and argued the same way: each (output row, input row) +// pair keeps its own accumulator, its own K-loop order, and its own simd_sum; +// `qdot_affine4_registered` is the bits == 4 arm of `qdot` verbatim. Only the +// LOADS are shared, so every output element's add sequence is identical to +// stock qmv_impl. +template +METAL_FUNC void qmv_affine4_g64_quad_stream_impl( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x0, + const device T* x1, + const device T* x2, + const device T* x3, + device T* y0, + device T* y1, + device T* y2, + device T* y3, + const int in_vec_size, + uint3 tid [[threadgroup_position_in_grid]], + uint simd_gid [[simdgroup_index_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + constexpr int num_simdgroups = 2; + constexpr int results_per_simdgroup = 4; + constexpr int values_per_thread = 8; + constexpr int block_size = values_per_thread * SIMD_SIZE; + constexpr int bytes_per_thread = 4; + constexpr int scale_step_per_thread = 8; + + const device uint8_t* ws = (const device uint8_t*)w; + thread float x_thread[values_per_thread]; + thread uint packed[results_per_simdgroup]; + thread float scale_local[results_per_simdgroup]; + thread float bias_local[results_per_simdgroup]; + thread float result0[results_per_simdgroup] = {0}; + thread float result1[results_per_simdgroup] = {0}; + thread float result2[results_per_simdgroup] = {0}; + thread float result3[results_per_simdgroup] = {0}; + + const int in_vec_size_w = in_vec_size / 2; + const int in_vec_size_g = in_vec_size / 64; + const int out_row = tid.y * (num_simdgroups * results_per_simdgroup) + + simd_gid * results_per_simdgroup; + + ws += out_row * in_vec_size_w + simd_lid * bytes_per_thread; + scales += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + biases += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + x0 += simd_lid * values_per_thread; + x1 += simd_lid * values_per_thread; + x2 += simd_lid * values_per_thread; + x3 += simd_lid * values_per_thread; + y0 += out_row; + y1 += out_row; + y2 += out_row; + y3 += out_row; + + int k = 0; + for (; k <= in_vec_size - block_size; k += block_size) { + for (int row = 0; row < results_per_simdgroup; row++) { + packed[row] = + *((const device uint*)(ws + row * in_vec_size_w)); + scale_local[row] = scales[row * in_vec_size_g]; + bias_local[row] = biases[row * in_vec_size_g]; + } + + float sum = load_vector(x0, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x1, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result1[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x2, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result2[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x3, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result3[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + + ws += block_size / 2; + scales += block_size / 64; + biases += block_size / 64; + x0 += block_size; + x1 += block_size; + x2 += block_size; + x3 += block_size; + } + + const int remaining = clamp( + static_cast(in_vec_size - k - simd_lid * values_per_thread), + 0, + values_per_thread); + if (remaining > 0) { + for (int row = 0; row < results_per_simdgroup; row++) { + packed[row] = + *((const device uint*)(ws + row * in_vec_size_w)); + scale_local[row] = scales[row * in_vec_size_g]; + bias_local[row] = biases[row * in_vec_size_g]; + } + + float sum = + load_vector_safe(x0, x_thread, remaining); + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = + load_vector_safe(x1, x_thread, remaining); + for (int row = 0; row < results_per_simdgroup; row++) { + result1[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = + load_vector_safe(x2, x_thread, remaining); + for (int row = 0; row < results_per_simdgroup; row++) { + result2[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = + load_vector_safe(x3, x_thread, remaining); + for (int row = 0; row < results_per_simdgroup; row++) { + result3[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + } + + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] = simd_sum(result0[row]); + result1[row] = simd_sum(result1[row]); + result2[row] = simd_sum(result2[row]); + result3[row] = simd_sum(result3[row]); + if (simd_lid == 0) { + y0[row] = static_cast(result0[row]); + y1[row] = static_cast(result1[row]); + y2[row] = static_cast(result2[row]); + y3[row] = static_cast(result3[row]); + } + } +} + +// Three-row weight-stream sharing: qmv_affine4_g64_quad_stream_impl with the +// fourth input row deleted, for same-expert gather runs of exactly three. The +// register discipline (one live x buffer), K-loop order, per-(output, input) +// accumulators, and qdot_affine4_registered arithmetic are the quad's own, so +// every output element's add sequence remains identical to stock qmv_impl. +template +METAL_FUNC void qmv_affine4_g64_triple_stream_impl( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x0, + const device T* x1, + const device T* x2, + device T* y0, + device T* y1, + device T* y2, + const int in_vec_size, + uint3 tid [[threadgroup_position_in_grid]], + uint simd_gid [[simdgroup_index_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + constexpr int num_simdgroups = 2; + constexpr int results_per_simdgroup = 4; + constexpr int values_per_thread = 8; + constexpr int block_size = values_per_thread * SIMD_SIZE; + constexpr int bytes_per_thread = 4; + constexpr int scale_step_per_thread = 8; + + const device uint8_t* ws = (const device uint8_t*)w; + thread float x_thread[values_per_thread]; + thread uint packed[results_per_simdgroup]; + thread float scale_local[results_per_simdgroup]; + thread float bias_local[results_per_simdgroup]; + thread float result0[results_per_simdgroup] = {0}; + thread float result1[results_per_simdgroup] = {0}; + thread float result2[results_per_simdgroup] = {0}; + + const int in_vec_size_w = in_vec_size / 2; + const int in_vec_size_g = in_vec_size / 64; + const int out_row = tid.y * (num_simdgroups * results_per_simdgroup) + + simd_gid * results_per_simdgroup; + + ws += out_row * in_vec_size_w + simd_lid * bytes_per_thread; + scales += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + biases += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + x0 += simd_lid * values_per_thread; + x1 += simd_lid * values_per_thread; + x2 += simd_lid * values_per_thread; + y0 += out_row; + y1 += out_row; + y2 += out_row; + + int k = 0; + for (; k <= in_vec_size - block_size; k += block_size) { + for (int row = 0; row < results_per_simdgroup; row++) { + packed[row] = + *((const device uint*)(ws + row * in_vec_size_w)); + scale_local[row] = scales[row * in_vec_size_g]; + bias_local[row] = biases[row * in_vec_size_g]; + } + + float sum = load_vector(x0, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x1, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result1[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x2, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result2[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + + ws += block_size / 2; + scales += block_size / 64; + biases += block_size / 64; + x0 += block_size; + x1 += block_size; + x2 += block_size; + } + + const int remaining = clamp( + static_cast(in_vec_size - k - simd_lid * values_per_thread), + 0, + values_per_thread); + if (remaining > 0) { + for (int row = 0; row < results_per_simdgroup; row++) { + packed[row] = + *((const device uint*)(ws + row * in_vec_size_w)); + scale_local[row] = scales[row * in_vec_size_g]; + bias_local[row] = biases[row * in_vec_size_g]; + } + + float sum = + load_vector_safe(x0, x_thread, remaining); + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = + load_vector_safe(x1, x_thread, remaining); + for (int row = 0; row < results_per_simdgroup; row++) { + result1[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = + load_vector_safe(x2, x_thread, remaining); + for (int row = 0; row < results_per_simdgroup; row++) { + result2[row] += qdot_affine4_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + } + + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] = simd_sum(result0[row]); + result1[row] = simd_sum(result1[row]); + result2[row] = simd_sum(result2[row]); + if (simd_lid == 0) { + y0[row] = static_cast(result0[row]); + y1[row] = static_cast(result1[row]); + y2[row] = static_cast(result2[row]); + } + } +} + +template +METAL_FUNC void qmv_affine8_g64_pair_impl( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x0, + const device T* x1, + device T* y0, + device T* y1, + const constant int& in_vec_size, + uint3 tid [[threadgroup_position_in_grid]], + uint simd_gid [[simdgroup_index_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + constexpr int num_simdgroups = 2; + constexpr int results_per_simdgroup = 4; + constexpr int values_per_thread = 4; + constexpr int block_size = values_per_thread * SIMD_SIZE; + constexpr int bytes_per_thread = 4; + constexpr int scale_step_per_thread = 16; + + const device uint8_t* ws = (const device uint8_t*)w; + thread float x0_thread[values_per_thread]; + thread float x1_thread[values_per_thread]; + thread float result0[results_per_simdgroup] = {0}; + thread float result1[results_per_simdgroup] = {0}; + + const int in_vec_size_w = in_vec_size; + const int in_vec_size_g = in_vec_size / 64; + const int out_row = tid.y * (num_simdgroups * results_per_simdgroup) + + simd_gid * results_per_simdgroup; + + ws += out_row * in_vec_size_w + simd_lid * bytes_per_thread; + scales += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + biases += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + x0 += simd_lid * values_per_thread; + x1 += simd_lid * values_per_thread; + y0 += out_row; + y1 += out_row; + + int k = 0; + for (; k <= in_vec_size - block_size; k += block_size) { + float sum0 = load_vector(x0, x0_thread); + float sum1 = load_vector(x1, x1_thread); + + for (int row = 0; row < results_per_simdgroup; row++) { + const device uint8_t* wl = ws + row * in_vec_size_w; + const device T* sl = scales + row * in_vec_size_g; + const device T* bl = biases + row * in_vec_size_g; + float dot0; + float dot1; + qdot_affine8_pair( + wl, x0_thread, x1_thread, sl[0], bl[0], sum0, sum1, dot0, dot1); + result0[row] += dot0; + result1[row] += dot1; + } + + ws += block_size; + scales += block_size / 64; + biases += block_size / 64; + x0 += block_size; + x1 += block_size; + } + + const int remaining = clamp( + static_cast(in_vec_size - k - simd_lid * values_per_thread), + 0, + values_per_thread); + if (remaining > 0) { + float sum0 = load_vector_safe( + x0, x0_thread, remaining); + float sum1 = load_vector_safe( + x1, x1_thread, remaining); + for (int row = 0; row < results_per_simdgroup; row++) { + const device uint8_t* wl = ws + row * in_vec_size_w; + const device T* sl = scales + row * in_vec_size_g; + const device T* bl = biases + row * in_vec_size_g; + float dot0; + float dot1; + qdot_affine8_pair( + wl, x0_thread, x1_thread, sl[0], bl[0], sum0, sum1, dot0, dot1); + result0[row] += dot0; + result1[row] += dot1; + } + } + + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] = simd_sum(result0[row]); + result1[row] = simd_sum(result1[row]); + if (simd_lid == 0) { + y0[row] = static_cast(result0[row]); + y1[row] = static_cast(result1[row]); + } + } +} + +// Four-row byte-weight-stream sharing for the affine8/g64 QMV, held to a +// PAIR-SIZED register footprint, exactly as +// qmv_affine4_g64_quad_stream_impl does for the nibble path. Same geometry as +// qmv_affine8_g64_pair_impl: two simdgroups by four output rows per 64-thread +// group, values_per_thread = 4, block_size = 128, scale_step_per_thread = 16, +// same load_vector / load_vector_safe tail. +// +// The block's byte weights and scale/bias are fetched into registers ONCE and +// the four cohort input rows then walk through a SINGLE four-value x buffer, +// so only one row of x is live at a time. This is the reason an earlier +// affine-8 quad that held all four rows of x live measured as a regression: +// the residency, not the arithmetic. Each (output row, input row) pair keeps +// its own accumulator, its own K-loop order and its own simd_sum, and +// qdot_affine8_registered is the bits == 8 arm of qdot verbatim, so every +// output element's add sequence is identical to stock qmv_impl -- only the +// LOADS are shared. +template +METAL_FUNC void qmv_affine8_g64_quad_stream_impl( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x0, + const device T* x1, + const device T* x2, + const device T* x3, + device T* y0, + device T* y1, + device T* y2, + device T* y3, + const constant int& in_vec_size, + uint3 tid [[threadgroup_position_in_grid]], + uint simd_gid [[simdgroup_index_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + constexpr int num_simdgroups = 2; + constexpr int results_per_simdgroup = 4; + constexpr int values_per_thread = 4; + constexpr int block_size = values_per_thread * SIMD_SIZE; + constexpr int bytes_per_thread = 4; + constexpr int scale_step_per_thread = 16; + + const device uint8_t* ws = (const device uint8_t*)w; + thread float x_thread[values_per_thread]; + thread uint packed[results_per_simdgroup]; + thread float scale_local[results_per_simdgroup]; + thread float bias_local[results_per_simdgroup]; + thread float result0[results_per_simdgroup] = {0}; + thread float result1[results_per_simdgroup] = {0}; + thread float result2[results_per_simdgroup] = {0}; + thread float result3[results_per_simdgroup] = {0}; + + const int in_vec_size_w = in_vec_size; + const int in_vec_size_g = in_vec_size / 64; + const int out_row = tid.y * (num_simdgroups * results_per_simdgroup) + + simd_gid * results_per_simdgroup; + + ws += out_row * in_vec_size_w + simd_lid * bytes_per_thread; + scales += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + biases += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + x0 += simd_lid * values_per_thread; + x1 += simd_lid * values_per_thread; + x2 += simd_lid * values_per_thread; + x3 += simd_lid * values_per_thread; + y0 += out_row; + y1 += out_row; + y2 += out_row; + y3 += out_row; + + int k = 0; + for (; k <= in_vec_size - block_size; k += block_size) { + for (int row = 0; row < results_per_simdgroup; row++) { + packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); + scale_local[row] = scales[row * in_vec_size_g]; + bias_local[row] = biases[row * in_vec_size_g]; + } + + float sum = load_vector(x0, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] += qdot_affine8_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x1, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result1[row] += qdot_affine8_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x2, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result2[row] += qdot_affine8_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x3, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result3[row] += qdot_affine8_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + + ws += block_size; + scales += block_size / 64; + biases += block_size / 64; + x0 += block_size; + x1 += block_size; + x2 += block_size; + x3 += block_size; + } + + // Dense Gemma 4 K is g64-aligned, so the tail is always a whole number of + // four-value lane packets. In particular down_proj K=2112 leaves exactly + // 16 active lanes; use the fixed unrolled load instead of four dynamic + // safe-tail loops while preserving each lane's qdot and simd_sum order. + const uint active_tail_lanes = + uint((in_vec_size - k) / values_per_thread); + if (simd_lid < active_tail_lanes) { + for (int row = 0; row < results_per_simdgroup; row++) { + packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); + scale_local[row] = scales[row * in_vec_size_g]; + bias_local[row] = biases[row * in_vec_size_g]; + } + + float sum = load_vector(x0, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] += qdot_affine8_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x1, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result1[row] += qdot_affine8_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x2, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result2[row] += qdot_affine8_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + sum = load_vector(x3, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + result3[row] += qdot_affine8_registered_word( + packed[row], x_thread, scale_local[row], bias_local[row], sum); + } + } + + for (int row = 0; row < results_per_simdgroup; row++) { + result0[row] = simd_sum(result0[row]); + result1[row] = simd_sum(result1[row]); + result2[row] = simd_sum(result2[row]); + result3[row] = simd_sum(result3[row]); + if (simd_lid == 0) { + y0[row] = static_cast(result0[row]); + y1[row] = static_cast(result1[row]); + y2[row] = static_cast(result2[row]); + y3[row] = static_cast(result3[row]); + } + } +} + +// GROUP-EXACT-MMA -- the fp32 `simdgroup_float8x8` body for the M = 8 decode +// cohort on 4-bit affine g64 weights. `A` holds the raw weight codes +// (8 output rows x 8 k-slots), `B` holds X^T (8 k-slots x 8 cohort rows) and +// `C` is zeroed per g64 group, so each group's 64 products are formed exactly +// (a bf16 x times a 4-bit code needs at most 12 significant bits) and summed +// by the matrix unit before the single combined `acc += s * C + rs * b` close +// in ascending k. Inside group g, fragment j and slot s name +// k(j, s) = 64 g + 8 s + j; a dot product is order free, so A and B may share +// any bijection, and this one makes every lane's fragment one contiguous +// load: 16 nibbles (`uint2`) of one weight row for A, two 8-value runs +// (`uint4`) of two cohort rows for B. +struct mma8_coord { + short fm; + short fn; +}; + +// steel/gemm/mma.h's `get_coord` arithmetic, reproduced locally so the same +// text compiles wherever this body is pasted: lane (fm, fn) owns elements +// (fm, fn) and (fm, fn + 1) of every 8x8 operand. +inline mma8_coord mma8_lane(uint lane) { + const short qid = short(lane / 4); + return { + short((qid & 4) + short((lane / 2) % 4)), + short((qid & 2) * 2 + short(lane % 2) * 2)}; +} + +// The x-side loads pull sixteen bytes at a time and split them into eight +// 16-bit lanes, which only makes sense for a 2-byte T. `affine_qmv` is also +// instantiated for `float`; the tier gate carries `sizeof(T) == 2` so the +// float instantiation never runs this body, and this primary template is what +// lets it still compile. +template +struct mma8_u16 { + static inline T cast(ushort u) { + return T(0); + } +}; + +template +struct mma8_u16 { + static inline T cast(ushort u) { + return as_type(u); + } +}; + +// Widening a 16-bit float to fp32 is exact, so these two reproduce the +// reference's own operand values bit for bit. +template +inline float mma8_lo(uint u) { + return float(mma8_u16::cast(ushort(u & 0xFFFFu))); +} + +template +inline float mma8_hi(uint u) { + return float(mma8_u16::cast(ushort(u >> 16))); +} - ws += block_size * bytes_per_pack / pack_factor; - scales += block_size / group_size; - biases += block_size / group_size; - x += block_size; - } - const int remaining = clamp( - static_cast(in_vec_size - k - simd_lid * values_per_thread), - 0, - values_per_thread); - if (remaining > 0) { - U sum = load_vector_safe( - x, x_thread, remaining); +// Textual twin of `load_vector`'s `sum` on the same aligned +// 8-run that the reference lane owns: the parenthesised 4-tuple is evaluated +// on T exactly as in the reference, then the two trees are added in fp32. The +// bias term of the affine form therefore reuses the reference's own +// elementary values, not a re-derived sum. +template +inline float mma8_runsum4(uint4 r) { + thread T xt[8]; + xt[0] = mma8_u16::cast(ushort(r.x & 0xFFFFu)); + xt[1] = mma8_u16::cast(ushort(r.x >> 16)); + xt[2] = mma8_u16::cast(ushort(r.y & 0xFFFFu)); + xt[3] = mma8_u16::cast(ushort(r.y >> 16)); + xt[4] = mma8_u16::cast(ushort(r.z & 0xFFFFu)); + xt[5] = mma8_u16::cast(ushort(r.z >> 16)); + xt[6] = mma8_u16::cast(ushort(r.w & 0xFFFFu)); + xt[7] = mma8_u16::cast(ushort(r.w >> 16)); + float sum = 0; + sum += xt[0] + xt[1] + xt[2] + xt[3]; + sum += xt[4] + xt[5] + xt[6] + xt[7]; + return sum; +} - for (int row = 0; row < results_per_simdgroup; row++) { - auto wl = (const device uint8_t*)(ws + row * in_vec_size_w); - const device T* sl = scales + row * in_vec_size_g; - const device T* bl = biases + row * in_vec_size_g; +#define MMA8_SETB(BB, W, HI) \ + BB.thread_elements()[0] = mma8_##HI(r0.W); \ + BB.thread_elements()[1] = mma8_##HI(r1.W); + +#define MMA8_STEP(BB, J) \ + A.thread_elements()[0] = float(extract_bits(wv.x, 4 * (J), 4)); \ + A.thread_elements()[1] = float(extract_bits(wv.y, 4 * (J), 4)); \ + simdgroup_multiply_accumulate(C, A, BB, C); + +// x is [8, K] with K % 64 == 0, w is packed [N, K / 8] uint32, scales and +// biases are [N, K / 64], y is [8, N]. `n0` is the first of the eight output +// rows this threadgroup owns. KS = 2 splits the K / 64 groups between the two +// simdgroups of the host's (32, 2, 1) threadgroup; an odd group count gives +// the extra group to simdgroup 0, which is deterministic and independent of +// scheduling. `red` is 32 float2 of threadgroup memory for the KS = 2 close. +template +METAL_FUNC void gemma4_qmv_mma8_affine4_g64_impl( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x, + device T* y, + const int K, + const int N, + const int n0, + threadgroup float2* red, + uint simd_gid, + uint simd_lid) { + const int G = K / 64; + const int gh = (G + 1) / 2; + const int g_begin = (KS == 2 && simd_gid == 1) ? gh : 0; + const int g_end = (KS == 2 && simd_gid == 0) ? gh : G; + const mma8_coord c = mma8_lane(simd_lid); + + const device uint8_t* wrow = + (const device uint8_t*)w + (n0 + c.fm) * (K / 2) + 4 * c.fn; + const device T* srow = scales + (n0 + c.fm) * G; + const device T* brow = biases + (n0 + c.fm) * G; + const device T* x0 = x + c.fn * K + 8 * c.fm; + const device T* x1 = x0 + K; + + float acc0 = 0.0f; + float acc1 = 0.0f; + simdgroup_float8x8 A; + simdgroup_float8x8 B0, B1, B2, B3, B4, B5, B6, B7; + + for (int g = g_begin; g < g_end; ++g) { + const uint4 r0 = *((const device uint4*)(x0 + 64 * g)); + const uint4 r1 = *((const device uint4*)(x1 + 64 * g)); + + // Each B lane owns the two 8-runs whose run sums the C lane (fm, fn) + // needs; three xor-butterfly steps over the fm lane bits broadcast + // RS[g][fn] and RS[g][fn + 1] to all eight lanes of the fn column group. + float2 rs = float2(mma8_runsum4(r0), mma8_runsum4(r1)); + rs += simd_shuffle_xor(rs, 2u); + rs += simd_shuffle_xor(rs, 4u); + rs += simd_shuffle_xor(rs, 16u); + + MMA8_SETB(B0, x, lo) + MMA8_SETB(B1, x, hi) + MMA8_SETB(B2, y, lo) + MMA8_SETB(B3, y, hi) + MMA8_SETB(B4, z, lo) + MMA8_SETB(B5, z, hi) + MMA8_SETB(B6, w, lo) + MMA8_SETB(B7, w, hi) + + const uint2 wv = *((const device uint2*)(wrow + 32 * g)); + const float s = float(srow[g]); + const float b = float(brow[g]); + + simdgroup_float8x8 C = simdgroup_float8x8(0.0f); + MMA8_STEP(B0, 0) + MMA8_STEP(B1, 1) + MMA8_STEP(B2, 2) + MMA8_STEP(B3, 3) + MMA8_STEP(B4, 4) + MMA8_STEP(B5, 5) + MMA8_STEP(B6, 6) + MMA8_STEP(B7, 7) + + acc0 += s * C.thread_elements()[0] + rs.x * b; + acc1 += s * C.thread_elements()[1] + rs.y * b; + } - U s = sl[0]; - U b = bl[0]; - result[row] += qdot_safe( - wl, x_thread, s, b, sum, remaining); - } + if (KS == 2) { + if (simd_gid == 1) { + red[simd_lid] = float2(acc0, acc1); } - for (int row = 0; row < results_per_simdgroup; row++) { - result[row] = simd_sum(result[row]); - if (simd_lid == 0) { - y[row] = static_cast(result[row]); - } + threadgroup_barrier(mem_flags::mem_threadgroup); + if (simd_gid == 1) { + return; } + const float2 other = red[simd_lid]; + acc0 = acc0 + other.x; + acc1 = acc1 + other.y; } + + y[c.fn * N + n0 + c.fm] = static_cast(acc0); + y[(c.fn + 1) * N + n0 + c.fm] = static_cast(acc1); } // Affine analog of fp_qmv_wide. Weights carry a scale and bias per group, so @@ -1750,6 +3170,7 @@ template < const constant int64_t* s_strides [[buffer(13)]], const constant int64_t* b_strides [[buffer(14)]], uint3 tid [[threadgroup_position_in_grid]], + uint3 ntg [[threadgroups_per_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], uint simd_lid [[thread_index_in_simdgroup]]) { if (batched) { @@ -1771,6 +3192,124 @@ template < b_strides, tid); } + if (!batched && group_size == 64 && bits == 2 && out_vec_size == 98336 && + ntg.x == 1) { + // M == 1 coarse draft readout (draft-rerank scheme): the ONE 2-bit shape + // in the scored path; proposal-only by construction (see kernel header). + qmv_fast_singlerow_affine2_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, tid, simd_gid, + simd_lid); + return; + } + if (!batched && group_size == 64 && bits == 4 && out_vec_size >= 1024) { + if (out_vec_size >= 4096) { + // Wide row sharing needs enough output tiles to keep the machine fed; + // below 4096 outputs the reduced x-group count thins the grid, so the + // promoted pair kernel is kept there byte-for-byte. + switch (ntg.x) { + case 2: + qmv_fast_crossrow_affine4_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 3: + qmv_fast_crossrow_affine4_g64_m( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 4: + qmv_fast_crossrow_affine4_g64_m( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 5: + qmv_fast_crossrow_affine4_g64_m( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 6: + qmv_fast_crossrow_affine4_g64_m( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 7: + qmv_fast_crossrow_affine4_g64_m( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 8: + // 3+3+2, not 4+4. M = 8 is the only hot width whose EVEN split needs + // two simultaneous vec accumulators in every active worker; + // M = 9 uses three-lane vectors and profiles CHEAPER despite more work + // (319 / 437 / 216 us for M = 7 / 8 / 9 in the public cross-row study) + // — a register cliff, not work scaling. + // Exact: these lanes carry INDEPENDENT input rows and are never reduced + // across (simd_sum reduces along K WITHIN a row), so moving a row from + // lane 3 of a four-wide vector to lane 0 of a two-wide one cannot + // reorder its scalar chain. Template admits it: M in [3,9], 8 % 3 == 2 + // (no one-row tail), IPG 3 inside the wide helper's [2,4]. + // Receipts: 85d5bca3 2.91143, yzxoi 2.92675. + // SYNERGY with the streak gate above, which is why they ship together: + // gate 2 reaches the width-8 verify SOONER, so this kernel fires MORE. + qmv_fast_crossrow_affine4_g64_m( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 9: + qmv_fast_crossrow_affine4_g64_m( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + default: + break; + } + } else { + switch (ntg.x) { + case 2: + qmv_fast_crossrow_affine4_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 3: + qmv_fast_crossrow_affine4_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 4: + qmv_fast_crossrow_affine4_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 5: + qmv_fast_crossrow_affine4_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 6: + qmv_fast_crossrow_affine4_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 7: + qmv_fast_crossrow_affine4_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 8: + qmv_fast_crossrow_affine4_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + case 9: + qmv_fast_crossrow_affine4_g64( + w, scales, biases, x, y, in_vec_size, out_vec_size, + tid, simd_gid, simd_lid); + return; + default: + break; + } + } + } qmv_fast_impl( w, scales, @@ -1808,6 +3347,7 @@ template < const constant int64_t* s_strides [[buffer(13)]], const constant int64_t* b_strides [[buffer(14)]], uint3 tid [[threadgroup_position_in_grid]], + uint3 ntg [[threadgroups_per_grid]], uint simd_gid [[simdgroup_index_in_threadgroup]], uint simd_lid [[thread_index_in_simdgroup]]) { if (batched) { @@ -1829,6 +3369,185 @@ template < b_strides, tid); } + if (!batched && group_size == 64 && bits == 4 && ntg.x == 8 && + ntg.z == 1 && in_vec_size % 64 == 0 && out_vec_size >= 8 && + out_vec_size % 8 == 0) { + // The ruled decode cohort presents eight input rows to ordinary QMV. + // MMA-QKV S1 -- GROUP-EXACT-MMA tier. It replaces, for the 4-bit affine + // g64 dense decode projections wide enough to fill the machine (q/k/v: + // N = 1024 / 2048 / 4096 / 8192 over K = 2816; the tied head only if its + // Swift MMA kernel is bypassed, since that road is tried first), the + // scalar quad_stream tier below: instead of eight per-lane 8-term chains + // that are each scaled and then reduced by `simd_sum`, one fp32 + // `simdgroup_float8x8` multiply-accumulate chain forms all 64 products of + // a g64 group and sums them before the single `s * C + rs * b` close. Every + // elementary term is the reference's own -- the products x * q are exact in + // fp32 (a bf16 x carries 8 significant bits, a code 4), scales and biases + // widen exactly, `mma8_runsum4` reproduces `load_vector`'s bf16 4-tuple sum + // order on the same aligned 8-run, and the group closes are chained in + // ascending k -- so the ONLY numeric deviation is fp32 reassociation inside + // the 64-wide group dot (plus the two-halves add of the KS = 2 split). This + // is the first non-bit-exact QMV tier here; measured against the stock M = 1 + // road over 50 random cohorts per plane at K = 2816, the deviation is at + // most 1 bf16 ulp for every output above the 2^-10 * row-max magnitude gate + // (non-zero fraction ~1.4e-4, 0 argmax flips over 400 rows per plane), it + // is run-to-run bitwise deterministic, and the body measured 0.41-0.51x the + // quad_stream body net of the dispatch floor on an M4 Max. Outputs cancelled + // below ~2^-8 of their term mass can show a second relative ulp; they are + // numerically negligible and never argmax candidates. KILL SWITCH: set + // `kGemma4QmvMma8Affine4` to false and this branch vanishes at compile time, + // restoring the quad_stream and pair tiers below byte for byte -- nothing + // beneath this block was edited. Raising `kGemma4QmvMma8Affine4FloorN` + // returns individual planes the same way. (MSL forbids a program-scope + // `constexpr`, so the two switches live at the top of the tier they guard.) + constexpr bool kGemma4QmvMma8Affine4 = true; + constexpr int kGemma4QmvMma8Affine4FloorN = 1024; + if (kGemma4QmvMma8Affine4 && sizeof(T) == 2 && ntg.z == 1 && + in_vec_size % 64 == 0 && + out_vec_size >= kGemma4QmvMma8Affine4FloorN && out_vec_size % 8 == 0) { + // Seven of the eight host x-groups retire before any load; the eighth + // produces all eight cohort columns of its eight output rows. + if (tid.x != 0) { + return; + } + threadgroup float2 red[32]; + gemma4_qmv_mma8_affine4_g64_impl( + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + 8 * int(tid.y), + red, + simd_gid, + simd_lid); + return; + } + if (out_vec_size >= 1024) { + // WIDE-N tier -- every non-`fast` 4-bit decode plane on this model: + // full-attention k_proj N = 1024 (k_eq_v), k/v_proj N = 2048, sliding + // q_proj N = 4096, full q_proj N = 8192, tied lm_head N = 262144. K is + // 2816, not a multiple of 512, so none of these reach affine_qmv_fast + // and its cross-row family; before this they ran two-row pair (and, at + // N >= 8192, a four-row quad that held 32 floats of x live and lost to + // the pair it replaced). + // + // One packed-weight stream feeds FOUR cohort rows in two active x-groups + // (4+4); the remaining host groups return. Per-row qdot, K-loop and + // simd_sum keep the stock qmv_impl sequence for every output element -- + // only loads are shared. + // + // Measured on the ranked box (M4 Pro, B = 8, streamed weight pool so + // every dispatch pulls from DRAM), us/dispatch, incumbent -> this: + // N = 1024 27.3 -> 26.1 N = 2048 54.1 -> 51.3 + // N = 4096 106.7 -> 101.6 N = 8192 233.6 -> 201.5 + // The floor sits at 1024 because that is the smallest plane measured to + // convert; below it (router.proj N = 128) two active x-groups leave only + // (N / 8) * 2 threadgroups and the promoted pair kernel is kept + // byte-for-byte. + const int first_m = int(tid.x) * 4; + if (first_m >= 8) { + return; + } + qmv_affine4_g64_quad_stream_impl( + w, + scales, + biases, + x + first_m * in_vec_size, + x + (first_m + 1) * in_vec_size, + x + (first_m + 2) * in_vec_size, + x + (first_m + 3) * in_vec_size, + y + first_m * out_vec_size, + y + (first_m + 1) * out_vec_size, + y + (first_m + 2) * out_vec_size, + y + (first_m + 3) * out_vec_size, + in_vec_size, + tid, + simd_gid, + simd_lid); + return; + } + // Claim adjacent rows in four active x-groups and let the remaining host + // groups return. The established pair helper shares each packed-weight + // load while preserving each row's qdot, K-loop, and simd_sum order. + const int first_m = int(tid.x) * 2; + if (first_m >= 8) { + return; + } + qmv_affine4_g64_pair_impl( + w, + scales, + biases, + x + first_m * in_vec_size, + x + (first_m + 1) * in_vec_size, + y + first_m * out_vec_size, + y + (first_m + 1) * out_vec_size, + in_vec_size, + tid, + simd_gid, + simd_lid); + return; + } + if (!batched && group_size == 64 && bits == 8 && ntg.x == 8 && + ntg.z == 1 && in_vec_size % 64 == 0 && out_vec_size >= 8 && + out_vec_size % 8 == 0) { + // Dense decode projections use byte weights. + if (out_vec_size >= 1024) { + // WIDE-N tier -- the dense MLP of all 30 layers: gate_proj and up_proj + // N = 2112 over K = 2816, down_proj N = 2816 over K = 2112. One + // byte-weight stream feeds FOUR cohort rows in two active x-groups + // (4+4); the remaining host groups return, and per-row qdot, K-loop and + // simd_sum stay the stock qmv_impl sequence. + // + // Measured on the ranked box (M4 Pro, B = 8, streamed weight pool), + // us/dispatch, incumbent pair -> this: + // N = 2112 (gate/up) 64.4 -> 56.2 N = 2816 (down) 67.4 -> 59.0 + // Same 1024 floor as the nibble tier, which keeps router.proj (N = 128) + // on the promoted pair kernel byte-for-byte. + const int first_m = int(tid.x) * 4; + if (first_m >= 8) { + return; + } + qmv_affine8_g64_quad_stream_impl( + w, + scales, + biases, + x + first_m * in_vec_size, + x + (first_m + 1) * in_vec_size, + x + (first_m + 2) * in_vec_size, + x + (first_m + 3) * in_vec_size, + y + first_m * out_vec_size, + y + (first_m + 1) * out_vec_size, + y + (first_m + 2) * out_vec_size, + y + (first_m + 3) * out_vec_size, + in_vec_size, + tid, + simd_gid, + simd_lid); + return; + } + // Pair adjacent cohort rows so each weight byte feeds both exact per-row + // dot-product streams. + const int first_m = int(tid.x) * 2; + if (first_m >= 8) { + return; + } + qmv_affine8_g64_pair_impl( + w, + scales, + biases, + x + first_m * in_vec_size, + x + (first_m + 1) * in_vec_size, + y + first_m * out_vec_size, + y + (first_m + 1) * out_vec_size, + in_vec_size, + tid, + simd_gid, + simd_lid); + return; + } qmv_impl( w, scales, @@ -2274,6 +3993,322 @@ template simd_lid); } +// One affine-4 dot product against a packed weight WORD (32 values / 8 +// nibbles) held in registers. Byte-for-byte the bits == 4 / +// values_per_thread == 8 arm of `qdot`: `ws[0]` is the low half-word and +// `ws[1]` the high half-word of the same aligned uint, so the four nibble +// masks, the two 4-term sums, the accumulation order over i and the +// `scale * accum + sum * bias` close are unchanged. Only the load shape +// differs -- one 4-byte load instead of two 2-byte loads. +inline float qdot_affine4_g64_word( + uint v, + const thread float* x_thread, + float scale, + float bias, + float sum) { + const uint lo = v & 0x0000FFFFu; + const uint hi = v >> 16; + float accum = 0; + accum += + (x_thread[0] * float(lo & 0x000fu) + x_thread[1] * float(lo & 0x00f0u) + + x_thread[2] * float(lo & 0x0f00u) + x_thread[3] * float(lo & 0xf000u)); + accum += + (x_thread[4] * float(hi & 0x000fu) + x_thread[5] * float(hi & 0x00f0u) + + x_thread[6] * float(hi & 0x0f00u) + x_thread[7] * float(hi & 0xf000u)); + return scale * accum + sum * bias; +} +// EXPERT-SINGLES: the SINGLETON arm of the routed-expert gather QMV. +// Diverse decode routing (8 streams x top-8 over 128 experts, sorted into +// 64 assignments) leaves most runs at length ONE, so the RUN-QUAD leader +// rule above hands the majority of both expert planes to the stock +// `qmv_impl` -- the one arm of the hot expert path that had never been +// microbenched (the dequant-once / prefetch / unroll knobs were only ever +// tried on the tied-head quad_stream body, where they lost). +// +// This is that arm with LOADS-ONLY rescheduling. Identical lane -> K +// mapping, identical per-block `load_vector` transform, identical +// eight-term `qdot` expression evaluated in the identical 4 + 4 grouping, +// identical per-row accumulator, identical `simd_sum` and store: every +// output element's add sequence is byte-for-byte the sequence `qmv_impl` +// produces for it. Only the SHAPE of the loads changes. +// +// WVEC : the two adjacent `uint16_t` loads the bits == 4 arm of `qdot` +// emits per (row, K-block) become ONE aligned 4-byte load. The +// packed row base is uint32-aligned at every block boundary +// (in_vec_size_w = K / 2 with K in {2816, 704}, lane offset +// simd_lid * 4, block stride 128), and `ws[0]` / `ws[1]` are the +// low / high half-words of that word, so the four nibble masks +// and their two 4-term sums are unchanged. +// PF : software prefetch of the NEXT block's four weight words. One +// x row is live in the singleton arm, so the +13..+40% extra +// live state that sank PF on the 4-row quad_stream body does not +// apply here. +// KFIX : in_vec_size as a compile-time constant. The gemma4 gate has +// already proven in_vec_size is 2816 (gate/up) or 704 (down), so +// the K-loop trip count and every stride fold constant-fold. +// +// Instantiated only under the gemma4 pair-geometry gate, which is +// compile-time false unless group_size == 64 && bits == 4; the affine-4 / +// g64 constants below are hardcoded exactly as `qmv_affine4_g64_pair_impl` +// hardcodes them. +template +METAL_FUNC void qmv_affine4_g64_singles_impl( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x, + device T* y, + const int in_vec_size_rt, + const int out_vec_size, + uint3 tid, + uint simd_gid, + uint simd_lid) { + constexpr int num_simdgroups = 2; + constexpr int results_per_simdgroup = 4; + constexpr int values_per_thread = 8; + constexpr int block_size = values_per_thread * SIMD_SIZE; + constexpr int bytes_per_thread = 4; + constexpr int scale_step_per_thread = 8; + constexpr int block_bytes = 128; + constexpr int qgroup = 64; + + const int in_vec_size = (KFIX > 0) ? KFIX : in_vec_size_rt; + + const device uint8_t* ws = (const device uint8_t*)w; + typedef float U; + thread U x_thread[values_per_thread]; + thread U result[results_per_simdgroup] = {0}; + + const int in_vec_size_w = in_vec_size / 2; + const int in_vec_size_g = in_vec_size / qgroup; + const int out_row = tid.y * (num_simdgroups * results_per_simdgroup) + + simd_gid * results_per_simdgroup; + const int used_out_row = min(out_vec_size - results_per_simdgroup, out_row); + if (out_row >= out_vec_size) { + return; + } + + ws += used_out_row * in_vec_size_w + simd_lid * bytes_per_thread; + scales += used_out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + biases += used_out_row * in_vec_size_g + simd_lid / scale_step_per_thread; + x += tid.x * in_vec_size + simd_lid * values_per_thread; + y += tid.x * out_vec_size + used_out_row; + + const int nblocks = in_vec_size / block_size; + const device uint8_t* ws0 = ws; + + thread uint wpf[results_per_simdgroup]; + if (PF) { + for (int row = 0; row < results_per_simdgroup; row++) { + wpf[row] = *((const device uint*)(ws0 + row * in_vec_size_w)); + } + } + + for (int blk = 0; blk < nblocks; blk++) { + U sum = load_vector(x, x_thread); + + thread uint wcur[results_per_simdgroup]; + if (PF) { + for (int row = 0; row < results_per_simdgroup; row++) { + wcur[row] = wpf[row]; + } + const int nextblk = (blk + 1 < nblocks) ? (blk + 1) : blk; + const device uint8_t* wsn = ws0 + nextblk * block_bytes; + for (int row = 0; row < results_per_simdgroup; row++) { + wpf[row] = *((const device uint*)(wsn + row * in_vec_size_w)); + } + } + + for (int row = 0; row < results_per_simdgroup; row++) { + const device T* sl = scales + row * in_vec_size_g; + const device T* bl = biases + row * in_vec_size_g; + U s = sl[0]; + U b = bl[0]; + if (PF) { + result[row] += qdot_affine4_g64_word(wcur[row], x_thread, s, b, sum); + } else if (WVEC) { + const uint v = *((const device uint*)(ws + row * in_vec_size_w)); + result[row] += qdot_affine4_g64_word(v, x_thread, s, b, sum); + } else { + auto wl = (const device uint8_t*)(ws + row * in_vec_size_w); + result[row] += qdot(wl, x_thread, s, b, sum); + } + } + + ws += block_bytes; + scales += block_size / qgroup; + biases += block_size / qgroup; + x += block_size; + } + + const int tail_values = in_vec_size - nblocks * block_size; + if (tail_values > 0) { + // Affine callers keep K a whole number of quantization groups and the + // block loop advances by whole blocks, so the tail is a whole number of + // values_per_thread lane packets (down_proj K = 704 leaves 192 = 24). + // The dynamic safe tail below is kept verbatim from `qmv_impl` for the + // genuinely partial packet no affine caller presents. + if (tail_values % values_per_thread != 0) { + const int remaining = clamp( + static_cast(tail_values - simd_lid * values_per_thread), + 0, + values_per_thread); + if (remaining > 0) { + U sum = load_vector_safe( + x, x_thread, remaining); + for (int row = 0; row < results_per_simdgroup; row++) { + auto wl = (const device uint8_t*)(ws + row * in_vec_size_w); + const device T* sl = scales + row * in_vec_size_g; + const device T* bl = biases + row * in_vec_size_g; + U s = sl[0]; + U b = bl[0]; + result[row] += qdot_safe( + wl, x_thread, s, b, sum, remaining); + } + } + } + const uint active_tail_lanes = uint(tail_values / values_per_thread); + if (tail_values % values_per_thread == 0 && simd_lid < active_tail_lanes) { + U sum = load_vector(x, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + const device T* sl = scales + row * in_vec_size_g; + const device T* bl = biases + row * in_vec_size_g; + U s = sl[0]; + U b = bl[0]; + if (WVEC || PF) { + const uint v = *((const device uint*)(ws + row * in_vec_size_w)); + result[row] += qdot_affine4_g64_word(v, x_thread, s, b, sum); + } else { + auto wl = (const device uint8_t*)(ws + row * in_vec_size_w); + result[row] += qdot(wl, x_thread, s, b, sum); + } + } + } + } + + for (int row = 0; row < results_per_simdgroup; row++) { + result[row] = simd_sum(result[row]); + if (simd_lid == 0) { + y[row] = static_cast(result[row]); + } + } +} + +// KERN-DOWN-TILE: y-tile coarsening for the K = 704 expert down gather +// (the only pair-geometry plane at that K; out_vec_size = 2816). The +// frozen host launches grid (1, N/8 = 352, 64), so every 64-thread group +// amortizes its serial run_offset scan, gather offset arithmetic and +// eight simd_sums over only ~3 K-blocks of stream (704 = 2 * 256 + 192) +// -- measured ~390 GB/s while the K = 2816 gate/up gathers move the same +// unique bytes at 479-589 GB/s. Here only every span-th y-group survives +// (the rest return before the scan); the survivor elects ONCE and then +// walks its span consecutive 8-row y-tiles serially through the verbatim +// pair impl -- or, for a pairless run position, the verbatim stock +// qmv_impl -- with tid.y rewritten to the tile index (a strip-walk +// pattern). Tile u is served by survivor (u / span) * span +// at loop step u % span and by no other group, so every output row keeps +// the IDENTICAL qdot sequence, accumulator, simd_sum and store the +// untiled arm produces for it: loads-only rescheduling, registers stay +// pair-sized. 352 divides by both spans, so no ragged tail. The pairless +// arm is tile-walked HERE because the stock fall-through derives out_row +// from tid.y inside qmv_impl -- follower tiles of a pairless assignment +// would otherwise never be written. Verified uint16-exact vs the +// per-assignment quantized_matmul oracle and vs the untiled arm at +// K = 704, N = 2816, 64 assignments over 128 experts, M = 8, spans 4 and +// 2, 3 seeds, NaN-filled outputs (parity-down-tile, 2026-08-28). +template +METAL_FUNC void gather_qmv_gemma4_down_tile( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x, + const device uint32_t* lhs_indices, + const device uint32_t* rhs_indices, + device T* y, + const constant int& in_vec_size, + const constant int& out_vec_size, + const uint lhs_stride, + const uint rhs_stride, + const int64_t x_stride, + const int64_t w_stride, + const int64_t s_stride, + const int64_t b_stride, + uint3 tid, + uint simd_gid, + uint simd_lid) { + constexpr int gemma4_down_tile_span = 4; // sweep alternate: 2 + if (tid.y % uint(gemma4_down_tile_span) != 0u) { + return; + } + const uint assignment = tid.z; + const uint32_t route_word = rhs_indices[assignment * rhs_stride]; + const bool expert_prefix_bounds = (route_word & 0x80000000u) != 0u; + const uint32_t expert = + expert_prefix_bounds ? (route_word & 0xffu) : route_word; + uint run_offset = 0; + if (expert_prefix_bounds) { + run_offset = (route_word >> 8) & 0x3fu; + } else { + for (uint prior = assignment; prior > 0; --prior) { + if (rhs_indices[(prior - 1) * rhs_stride] != expert) { + break; + } + run_offset++; + } + } + // Odd positions are produced by the immediately preceding pair leader. + if ((run_offset & 1) != 0) { + return; + } + const device uint32_t* tile_w = w + expert * w_stride; + const device T* tile_scales = scales + expert * s_stride; + const device T* tile_biases = biases + expert * b_stride; + const device T* tile_x0 = + x + lhs_indices[assignment * lhs_stride] * x_stride; + device T* tile_y0 = y + assignment * out_vec_size; + const bool has_pair = expert_prefix_bounds + ? (((route_word >> 14) & 0x3fu) + 1u) > 1u + : assignment + 1 < 64 && + rhs_indices[(assignment + 1) * rhs_stride] == expert; + if (has_pair) { + const device T* tile_x1 = + x + lhs_indices[(assignment + 1) * lhs_stride] * x_stride; + device T* tile_y1 = y + (assignment + 1) * out_vec_size; + for (int t = 0; t < gemma4_down_tile_span; t++) { + uint3 tile_tid = tid; + tile_tid.y = tid.y + uint(t); + qmv_affine4_g64_pair_impl( + tile_w, + tile_scales, + tile_biases, + tile_x0, + tile_x1, + tile_y0, + tile_y1, + in_vec_size, + tile_tid, + simd_gid, + simd_lid); + } + return; + } + for (int t = 0; t < gemma4_down_tile_span; t++) { + uint3 tile_tid = tid; + tile_tid.y = tid.y + uint(t); + qmv_impl( + tile_w, + tile_scales, + tile_biases, + tile_x0, + tile_y0, + in_vec_size, + out_vec_size, + tile_tid, + simd_gid, + simd_lid); + } +} + template [[kernel]] void affine_gather_qmv( const device uint32_t* w [[buffer(0)]], @@ -2301,28 +4336,219 @@ template uint simd_gid [[simdgroup_index_in_threadgroup]], uint simd_lid [[thread_index_in_simdgroup]]) { int M = x_shape[x_batch_ndims]; - adjust_matrix_offsets( - x, - w, - scales, - biases, - lhs_indices, - rhs_indices, - y, - out_vec_size * M, - batch_ndims, - batch_shape, - lhs_strides, - rhs_strides, - x_batch_ndims, - x_shape, - x_strides, - w_batch_ndims, - w_shape, - w_strides, - s_strides, - b_strides, - tid); + const bool gemma4_pair_geometry = + group_size == 64 && bits == 4 && M == 1 && batch_ndims == 1 && + batch_shape[0] == 64 && x_batch_ndims == 1 && w_batch_ndims == 1 && + ((in_vec_size == 2816 && out_vec_size == 704) || + (in_vec_size == 704 && out_vec_size == 2816)); + if (gemma4_pair_geometry) { + // KERN-DOWN-TILE gate (strip-walk pattern): compile-time flip; ON + // here -- the K = 704 down plane takes the y-tile-coarsened arm above. + // Flip to false to return every plane to the incumbent per-y-group + // election below; the two arms are bit-identical by construction. + constexpr bool gemma4_down_tile = true; + if (gemma4_down_tile && in_vec_size == 704) { + gather_qmv_gemma4_down_tile( + w, + scales, + biases, + x, + lhs_indices, + rhs_indices, + y, + in_vec_size, + out_vec_size, + (uint)lhs_strides[0], + (uint)rhs_strides[0], + x_strides[0], + w_strides[0], + s_strides[0], + b_strides[0], + tid, + simd_gid, + simd_lid); + return; + } + const uint assignment = tid.z; + const uint32_t route_word = + rhs_indices[assignment * (uint)rhs_strides[0]]; + const bool expert_prefix_bounds = (route_word & 0x80000000u) != 0u; + const uint32_t expert = + expert_prefix_bounds ? (route_word & 0xffu) : route_word; + uint run_offset = 0; + if (expert_prefix_bounds) { + run_offset = (route_word >> 8) & 0x3fu; + } else { + for (uint prior = assignment; prior > 0; --prior) { + if (rhs_indices[(prior - 1) * (uint)rhs_strides[0]] != expert) { + break; + } + run_offset++; + } + } + + // RUN-QUAD: leaders sit at run_offset % 4 == 0 and serve up to four + // same-expert assignments from ONE weight stream. Positions 1..3 of each + // aligned quartet are produced by their leader, so a run of two keeps the + // incumbent pair arithmetic, a run of three takes the triple impl, and a + // run of four takes the quad-stream impl -- each (output, input) pair + // keeps its own accumulator, K-loop order, and qdot, so every output + // element's add sequence is identical to the incumbent per-arm kernels. + if ((run_offset & 3) != 0) { + return; + } + uint run_len = 1; + if (expert_prefix_bounds) { + run_len = min(4u, ((route_word >> 14) & 0x3fu) + 1u); + } else { + while (run_len < 4 && assignment + run_len < 64 && + rhs_indices[(assignment + run_len) * (uint)rhs_strides[0]] == + expert) { + run_len++; + } + } + if (run_len > 1) { + const device uint32_t* run_w = w + expert * w_strides[0]; + const device T* run_scales = scales + expert * s_strides[0]; + const device T* run_biases = biases + expert * b_strides[0]; + const device T* run_x0 = + x + lhs_indices[assignment * (uint)lhs_strides[0]] * x_strides[0]; + const device T* run_x1 = x + + lhs_indices[(assignment + 1) * (uint)lhs_strides[0]] * x_strides[0]; + device T* run_y0 = y + assignment * out_vec_size; + device T* run_y1 = y + (assignment + 1) * out_vec_size; + if (run_len == 2) { + qmv_affine4_g64_pair_impl( + run_w, + run_scales, + run_biases, + run_x0, + run_x1, + run_y0, + run_y1, + in_vec_size, + tid, + simd_gid, + simd_lid); + return; + } + const device T* run_x2 = x + + lhs_indices[(assignment + 2) * (uint)lhs_strides[0]] * x_strides[0]; + device T* run_y2 = y + (assignment + 2) * out_vec_size; + if (run_len == 3) { + qmv_affine4_g64_triple_stream_impl( + run_w, + run_scales, + run_biases, + run_x0, + run_x1, + run_x2, + run_y0, + run_y1, + run_y2, + in_vec_size, + tid, + simd_gid, + simd_lid); + return; + } + const device T* run_x3 = x + + lhs_indices[(assignment + 3) * (uint)lhs_strides[0]] * x_strides[0]; + device T* run_y3 = y + (assignment + 3) * out_vec_size; + qmv_affine4_g64_quad_stream_impl( + run_w, + run_scales, + run_biases, + run_x0, + run_x1, + run_x2, + run_x3, + run_y0, + run_y1, + run_y2, + run_y3, + in_vec_size, + tid, + simd_gid, + simd_lid); + return; + } + + // Singleton experts and odd-run tails still need one ordinary QMV, but + // their one-dimensional offsets are already resolved by this guard. + const uint32_t single_lhs = + lhs_indices[assignment * (uint)lhs_strides[0]]; + const device T* single_x = x + single_lhs * x_strides[0]; + const device uint32_t* single_w = w + expert * w_strides[0]; + const device T* single_scales = scales + expert * s_strides[0]; + const device T* single_biases = biases + expert * b_strides[0]; + device T* single_y = y + assignment * (uint)out_vec_size; + if (in_vec_size == 2816) { + qmv_affine4_g64_singles_impl< + T, group_size, bits, 2816, true, false>( + single_w, single_scales, single_biases, single_x, single_y, + in_vec_size, out_vec_size, tid, simd_gid, simd_lid); + } else { + qmv_impl( + single_w, single_scales, single_biases, single_x, single_y, + in_vec_size, out_vec_size, tid, simd_gid, simd_lid); + } + return; + } + uint32_t x_idx; + uint32_t route_word; + if (batch_ndims == 1) { + x_idx = lhs_indices[tid.z * lhs_strides[0]]; + route_word = rhs_indices[tid.z * rhs_strides[0]]; + } else { + ulong2 idx = elem_to_loc_broadcast( + tid.z, batch_shape, lhs_strides, rhs_strides, batch_ndims); + x_idx = lhs_indices[idx.x]; + route_word = rhs_indices[idx.y]; + } + if ((route_word & 0x80000000u) != 0u) { + const uint32_t expert = route_word & 0xffu; + if (x_batch_ndims == 1) { + x += x_idx * x_strides[0]; + } else { + x += elem_to_loc(x_idx, x_shape, x_strides, x_batch_ndims); + } + if (w_batch_ndims == 1) { + w += expert * w_strides[0]; + scales += expert * s_strides[0]; + biases += expert * b_strides[0]; + } else { + ulong3 idx = elem_to_loc_broadcast( + expert, w_shape, w_strides, s_strides, b_strides, w_batch_ndims); + w += idx.x; + scales += idx.y; + biases += idx.z; + } + y += tid.z * (out_vec_size * M); + } else { + adjust_matrix_offsets( + x, + w, + scales, + biases, + lhs_indices, + rhs_indices, + y, + out_vec_size * M, + batch_ndims, + batch_shape, + lhs_strides, + rhs_strides, + x_batch_ndims, + x_shape, + x_strides, + w_batch_ndims, + w_shape, + w_strides, + s_strides, + b_strides, + tid); + } qmv_impl( w, scales, From af74c88ba3c4fe3d0e5a3be004c0445464c138e0 Mon Sep 17 00:00:00 2001 From: David Tai <8346495+davidtai@users.noreply.github.com> Date: Thu, 3 Sep 2026 18:17:31 -0500 Subject: [PATCH 05/11] perf(metal): a D=512 two-pass vector SDPA cell for Gemma 4 decode 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 Claude-Session: https://claude.ai/code/session_014L39jXT8ReGzjxUoUKLfan --- .../scaled_dot_product_attention.metal | 23 ++++- .../metal/scaled_dot_product_attention.cpp | 84 ++++++++++++++++++- 2 files changed, 103 insertions(+), 4 deletions(-) diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal index 187e6fffca..27db133e17 100644 --- a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal @@ -38,6 +38,24 @@ using namespace metal; 8, \ hpt) +// D512-2PASS. The 2-pass split-K kernels ONLY, with no single-pass twin. +// +// `sdpa_vector` holds q, k and o in registers at D / 32 floats each and is +// launched at the Metal maximum 1024 threads per threadgroup +// (scaled_dot_product_attention.cpp:461). At D = 512 that is 48 live floats +// per thread against 1024 threads, an occupancy claim the 2-pass kernel does +// not make (it runs 32 x gqa_factor threads with 32 live floats), and the +// host routes every D = 512 call to the 2-pass form for exactly that reason. +// Instantiating the single-pass twin would only add a pipeline nothing can +// dispatch. +#define instantiate_sdpa_vector_2pass(type, qk_dim, value_dim) \ + instantiate_kernel( \ + "sdpa_vector_2pass_1_" #type "_" #qk_dim "_" #value_dim, \ + sdpa_vector_2pass_1, \ + type, \ + qk_dim, \ + value_dim) + #define instantiate_sdpa_vector_heads(type) \ instantiate_sdpa_vector(type, 64, 64) \ instantiate_sdpa_vector(type, 96, 96) \ @@ -45,13 +63,16 @@ using namespace metal; instantiate_sdpa_vector(type, 192, 128) \ instantiate_sdpa_vector(type, 192, 192) \ instantiate_sdpa_vector(type, 256, 256) \ + instantiate_sdpa_vector_2pass(type, 512, 512) \ instantiate_sdpa_vector_gqa(type, 64, 64, 8) \ instantiate_sdpa_vector_gqa(type, 128, 128, 4) \ + instantiate_sdpa_vector_gqa(type, 512, 512, 2) \ instantiate_sdpa_vector_aggregation(type, 64) \ instantiate_sdpa_vector_aggregation(type, 96) \ instantiate_sdpa_vector_aggregation(type, 128) \ instantiate_sdpa_vector_aggregation(type, 192) \ - instantiate_sdpa_vector_aggregation(type, 256) + instantiate_sdpa_vector_aggregation(type, 256) \ + instantiate_sdpa_vector_aggregation(type, 512) instantiate_sdpa_vector_heads(float) instantiate_sdpa_vector_heads(bfloat16_t) diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index 981d2b1817..f2faf2caa9 100644 --- a/mlx/backend/metal/scaled_dot_product_attention.cpp +++ b/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -1,5 +1,8 @@ // Copyright © 2024 Apple Inc. +#include +#include #include +#include #include "mlx/backend/common/compiled.h" #include "mlx/backend/gpu/copy.h" @@ -15,6 +18,71 @@ namespace mlx::core::fast { namespace { +// D512-2PASS (C1 rung 2). Gemma 4's five global attention layers are +// 16 query heads / 2 KV heads at head_dim 512, so `has_fused_kernel`'s +// vector head-dim list ({64, 96, 128, 192, 256}) rejects every decode call +// and `use_fallback` sends them to the unfused +// `scale.q -> matmul -> softmax -> matmul` graph, whose two matmuls are +// batched gemvs over the GQA-expanded batch: each of the 2 key planes and +// each of the 2 value planes is streamed once PER QUERY HEAD, eight times. +// `sdpa_vector_2pass_1` / `_2` are templated on D and V and need no new +// body at 512; this admits the dim and instantiates them +// (kernels/scaled_dot_product_attention.metal). +// +// 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. +// +// Off value: `DARKBLOOM_GEMMA4_D512_DECODE_2PASS` in {0, false, no, off} +// removes the admission and restores the unfused graph exactly. Default ON. +inline bool env_flag_on(const char* name, bool default_on) { + const char* raw = std::getenv(name); + if (raw == nullptr) { + return default_on; + } + std::string value(raw); + for (auto& c : value) { + c = static_cast(std::tolower(static_cast(c))); + } + if (value == "0" || value == "false" || value == "no" || value == "off") { + return false; + } + return true; +} + +inline bool d512_vector_sdpa_enabled() { + static bool enabled = + env_flag_on("DARKBLOOM_GEMMA4_D512_DECODE_2PASS", true); + return enabled; +} + +// D512-2PASS-DEDUP. `sdpa_vector_2pass_1` gives every query head of a GQA +// group its own simdgroup, so it still reads each K/V byte `gqa_factor` +// times -- it removes the dispatch chain and the materialised score plane, +// not the redundant stream. `sdpa_vector_2pass_1_gqa` is the +// duplication-free variant already in the tree (instantiated at 64/HPT=8 +// and 128/HPT=4): each simdgroup owns a token sub-chunk and carries HPT +// query heads, so each byte is read gqa_factor / HPT times. +// +// At D = 512 a thread holds `HPT * (D / 32)` query floats and the same +// number of output floats, so HPT = 2 (64 live floats, twice the plain +// 2-pass kernel's 32) is the only step that is clearly affordable; it +// halves the K/V stream. It is DEFAULT OFF because the register claim at +// 32 x gqa_factor = 256 threads per threadgroup has not been checked +// against `maxTotalThreadsPerThreadgroup` on the device, and a pipeline +// that comes back under 256 would make `check_kernel_threadgroup_size` +// throw rather than degrade. Turn it on with +// `DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP=1` once that is measured. +inline bool d512_gqa_dedup_enabled() { + static bool enabled = + env_flag_on("DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP", false); + return enabled; +} + +// The one head dim this port adds to the vector path. +constexpr int kD512 = 512; + void sdpa_full_self_attention_nax( const Stream& s, metal::Device& d, @@ -467,7 +535,9 @@ void sdpa_vector_2pass( kname.reserve(64); kname += "sdpa_vector_2pass_1"; if (!mask && !sinks && q.shape(2) == 1 && q.shape(1) == 8 * k.shape(1) && - q.shape(-1) == v.shape(-1) && (q.shape(-1) == 64 || q.shape(-1) == 128) && + q.shape(-1) == v.shape(-1) && + (q.shape(-1) == 64 || q.shape(-1) == 128 || + (q.shape(-1) == kD512 && d512_gqa_dedup_enabled())) && k.shape(2) >= 8192) { kname += "_gqa"; } @@ -683,7 +753,8 @@ std::tuple has_fused_kernel( (query_head_dim == value_head_dim && (query_head_dim == 64 || query_head_dim == 96 || query_head_dim == 128 || query_head_dim == 192 || - query_head_dim == 256)) || + query_head_dim == 256 || + (query_head_dim == kD512 && d512_vector_sdpa_enabled()))) || (query_head_dim == 192 && value_head_dim == 128); if (!supported_head_dim) { msg << "the vector attention kernel supports head dims " @@ -872,7 +943,14 @@ void ScaledDotProductAttention::eval_gpu( // - The sequence length is even longer and we have gqa bool do_causal = do_causal_ && q.shape(2) > 1; char devc = d.get_architecture().back(); - if (((devc == 'd' || devc == 's') && k.shape(2) >= 1024) || + // D512-2PASS: head dim 512 has no single-pass instantiation (see + // kernels/scaled_dot_product_attention.metal), so it takes the split-K + // form at EVERY key length, not only past the device thresholds below. + // The 2-pass kernel is length-generic: blocks with no keys leave + // `sums = 0` / `maxs = finite_min`, which the merge pass folds in with + // weight `exp(finite_min - max) == 0`. + if (q.shape(-1) == kD512 || + ((devc == 'd' || devc == 's') && k.shape(2) >= 1024) || (k.shape(1) < q.shape(1) && k.shape(2) >= 4096)) { sdpa_vector_2pass(s, d, q, k, v, o, scale_, do_causal, mask, sinks); } else { From bb794a7edf96bdd91ce4ee77f4fa4c18be360ff1 Mon Sep 17 00:00:00 2001 From: David Tai <8346495+davidtai@users.noreply.github.com> Date: Thu, 3 Sep 2026 19:28:41 -0500 Subject: [PATCH 06/11] fix(metal): fit the D512 dedup merge plane inside the 32 KB threadgroup 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`. 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 --- .../scaled_dot_product_attention.metal | 15 ++- mlx/backend/metal/kernels/sdpa_vector.h | 94 ++++++++++++++----- .../metal/scaled_dot_product_attention.cpp | 25 +++-- 3 files changed, 101 insertions(+), 33 deletions(-) diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal index 27db133e17..839d697027 100644 --- a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal @@ -28,7 +28,11 @@ using namespace metal; qk_dim, \ value_dim) -#define instantiate_sdpa_vector_gqa(type, qk_dim, value_dim, hpt) \ +// `split` is the merge-plane publish count, NOT part of the kernel name -- +// the host names this kernel by type and dims only. SPLIT = 1 is the shipped +// body instruction for instruction; a larger value only shrinks the +// threadgroup allocation (see sdpa_vector.h). +#define instantiate_sdpa_vector_gqa(type, qk_dim, value_dim, hpt, split) \ instantiate_kernel( \ "sdpa_vector_2pass_1_gqa_" #type "_" #qk_dim "_" #value_dim, \ sdpa_vector_2pass_1_gqa, \ @@ -36,7 +40,8 @@ using namespace metal; qk_dim, \ value_dim, \ 8, \ - hpt) + hpt, \ + split) // D512-2PASS. The 2-pass split-K kernels ONLY, with no single-pass twin. // @@ -64,9 +69,9 @@ using namespace metal; instantiate_sdpa_vector(type, 192, 192) \ instantiate_sdpa_vector(type, 256, 256) \ instantiate_sdpa_vector_2pass(type, 512, 512) \ - instantiate_sdpa_vector_gqa(type, 64, 64, 8) \ - instantiate_sdpa_vector_gqa(type, 128, 128, 4) \ - instantiate_sdpa_vector_gqa(type, 512, 512, 2) \ + instantiate_sdpa_vector_gqa(type, 64, 64, 8, 1) \ + instantiate_sdpa_vector_gqa(type, 128, 128, 4, 1) \ + instantiate_sdpa_vector_gqa(type, 512, 512, 2, 2) \ instantiate_sdpa_vector_aggregation(type, 64) \ instantiate_sdpa_vector_aggregation(type, 96) \ instantiate_sdpa_vector_aggregation(type, 128) \ diff --git a/mlx/backend/metal/kernels/sdpa_vector.h b/mlx/backend/metal/kernels/sdpa_vector.h index 23194a3f00..fc5731476b 100644 --- a/mlx/backend/metal/kernels/sdpa_vector.h +++ b/mlx/backend/metal/kernels/sdpa_vector.h @@ -333,7 +333,24 @@ template // each K/V byte is read G / HPT times instead of G times. Single-token // queries without mask or sinks only; the partials layout matches // sdpa_vector_2pass_2. -template +// +// SPLIT is how many passes the cross-simdgroup merge plane is published in. +// The plane is the kernel's only large threadgroup allocation and is +// G * HPT * V floats at SPLIT = 1 -- 16,640 B at (64, HPT 8) and +// (128, HPT 4), but 32,896 B at (512, HPT 2), which is 128 B over Metal's +// 32,768 B threadgroup limit and makes pipeline creation fail outright. +// SPLIT = n publishes V / n columns at a time and so allocates +// G * HPT * V / n floats. +// +// SPLIT does not change the arithmetic. Each lane keeps the same register +// slice, the shared plane is only a scratch relabelling of that slice (write +// and read use the identical lane mapping, never the global column index), +// `gmax` and `denom` are computed once from the full scalar arrays, and the +// per-output accumulation over s keeps its order inside every pass. SPLIT = 1 +// is the shipped code path instruction for instruction: one publish of the +// plane AND the scalars, one barrier, one merge -- the extra write-after-read +// barrier only exists for p > 0. +template [[kernel]] void sdpa_vector_2pass_1_gqa( const device T* queries [[buffer(0)]], const device T* keys [[buffer(1)]], @@ -429,33 +446,66 @@ template } } - threadgroup U o_sh[G * HPT * V]; + constexpr int VS = V / SPLIT; + constexpr int v_per_pass = v_per_thread / SPLIT; + + static_assert(V % SPLIT == 0, "sdpa_vector_2pass_1_gqa: SPLIT must divide V"); + static_assert( + v_per_thread % SPLIT == 0, + "sdpa_vector_2pass_1_gqa: SPLIT must divide V / 32"); + // The failure this guards is not hypothetical: (D 512, HPT 2) at SPLIT = 1 + // allocates 32,896 B and Metal refuses the pipeline at LOAD time with + // "Threadgroup memory size (32896) exceeds the maximum threadgroup memory + // allowed (32768)" -- a runtime error on the device, far from the code that + // caused it. Raising SPLIT is the fix; this turns forgetting to into a + // compile error at the offline `xcrun metal -c` gate. + static_assert( + (G * HPT * VS + 2 * G * HPT) * sizeof(U) <= 32768, + "sdpa_vector_2pass_1_gqa: the merge plane exceeds Metal's 32 KB " + "threadgroup limit -- raise SPLIT for this instantiation"); + + threadgroup U o_sh[G * HPT * VS]; threadgroup U se_sh[G * HPT]; threadgroup U mx_sh[G * HPT]; - for (int j = 0; j < HPT; j++) { - int slot = (h0 + j) * HPT + cchunk; - U inv = sum_exp_score[j] > 0 ? 1 / sum_exp_score[j] : 0; - for (int i = 0; i < v_per_thread; i++) { - o_sh[slot * V + simd_lid * v_per_thread + i] = o[j][i] * inv; - } - if (simd_lid == 0) { - se_sh[slot] = sum_exp_score[j]; - mx_sh[slot] = max_score[j]; - } - } - threadgroup_barrier(mem_flags::mem_threadgroup); U gmax = Limits::finite_min; - for (int s = 0; s < HPT; s++) { - gmax = max(gmax, mx_sh[g * HPT + s]); - } U denom = 0; U acc[v_per_thread] = {0}; - for (int s = 0; s < HPT; s++) { - U w = se_sh[g * HPT + s] * fast::exp(mx_sh[g * HPT + s] - gmax); - denom += w; - for (int i = 0; i < v_per_thread; i++) { - acc[i] += w * o_sh[(g * HPT + s) * V + simd_lid * v_per_thread + i]; + + for (int p = 0; p < SPLIT; p++) { + if (p > 0) { + // Write-after-read: every lane must finish reading the previous + // pass's plane before any lane overwrites it. + threadgroup_barrier(mem_flags::mem_threadgroup); + } + for (int j = 0; j < HPT; j++) { + int slot = (h0 + j) * HPT + cchunk; + U inv = sum_exp_score[j] > 0 ? 1 / sum_exp_score[j] : 0; + for (int i = 0; i < v_per_pass; i++) { + o_sh[slot * VS + simd_lid * v_per_pass + i] = + o[j][p * v_per_pass + i] * inv; + } + if (p == 0 && simd_lid == 0) { + se_sh[slot] = sum_exp_score[j]; + mx_sh[slot] = max_score[j]; + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + if (p == 0) { + for (int s = 0; s < HPT; s++) { + gmax = max(gmax, mx_sh[g * HPT + s]); + } + } + for (int s = 0; s < HPT; s++) { + U w = se_sh[g * HPT + s] * fast::exp(mx_sh[g * HPT + s] - gmax); + if (p == 0) { + denom += w; + } + for (int i = 0; i < v_per_pass; i++) { + acc[p * v_per_pass + i] += + w * o_sh[(g * HPT + s) * VS + simd_lid * v_per_pass + i]; + } } } diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index f2faf2caa9..cdc53aafcd 100644 --- a/mlx/backend/metal/scaled_dot_product_attention.cpp +++ b/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -68,12 +68,25 @@ inline bool d512_vector_sdpa_enabled() { // At D = 512 a thread holds `HPT * (D / 32)` query floats and the same // number of output floats, so HPT = 2 (64 live floats, twice the plain // 2-pass kernel's 32) is the only step that is clearly affordable; it -// halves the K/V stream. It is DEFAULT OFF because the register claim at -// 32 x gqa_factor = 256 threads per threadgroup has not been checked -// against `maxTotalThreadsPerThreadgroup` on the device, and a pipeline -// that comes back under 256 would make `check_kernel_threadgroup_size` -// throw rather than degrade. Turn it on with -// `DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP=1` once that is measured. +// halves the K/V stream. +// +// First device run failed pipeline creation outright: +// Threadgroup memory size (32896) exceeds the maximum threadgroup memory +// allowed (32768) +// -- the merge plane `o_sh[G * HPT * V]` is 8 * 2 * 512 floats = 32,768 B on +// its own, and the two 16-float scalar arrays put it 128 B over. Note that +// `blocks` is not a term in that expression, so tuning MLX_SDPA_BLOCKS could +// not have helped. Fixed by publishing the plane in SPLIT = 2 passes +// (kernels/sdpa_vector.h), which allocates 16,512 B -- in line with the +// shipped 64/128 instantiations' 16,640 B -- and leaves the arithmetic +// unchanged. +// +// Still DEFAULT OFF: the remaining unmeasured claim is the REGISTER one. +// A thread holds q[2][16] + o[2][16] + kr[16] + vr[16] + acc[16] = 112 live +// floats at 32 x gqa_factor = 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 with +// `DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP=1` for its own arm. inline bool d512_gqa_dedup_enabled() { static bool enabled = env_flag_on("DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP", false); From eb01609a690d3957a31dba897cf8e7fbb0c86d98 Mon Sep 17 00:00:00 2001 From: David Tai Date: Sun, 4 Oct 2026 20:54:20 -0500 Subject: [PATCH 07/11] fix(metal): name the D512 2-pass kernel as the host builds it Main #16 renamed the 2-pass vector SDPA kernels to sdpa_vector_2pass_fp32partials_*. The host builds "sdpa_vector_2pass_fp32partials_1__512_512", but the D512 instantiation still emitted "sdpa_vector_2pass_1__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 --- .../scaled_dot_product_attention.metal | 16 +++--- python/tests/test_fast_sdpa.py | 51 +++++++++++++++++++ 2 files changed, 59 insertions(+), 8 deletions(-) diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal index a0fbb0956b..62c71c5935 100644 --- a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal @@ -46,16 +46,16 @@ using namespace metal; // D512-2PASS. The 2-pass split-K kernels ONLY, with no single-pass twin. // // `sdpa_vector` holds q, k and o in registers at D / 32 floats each and is -// launched at the Metal maximum 1024 threads per threadgroup -// (scaled_dot_product_attention.cpp:461). At D = 512 that is 48 live floats -// per thread against 1024 threads, an occupancy claim the 2-pass kernel does -// not make (it runs 32 x gqa_factor threads with 32 live floats), and the -// host routes every D = 512 call to the 2-pass form for exactly that reason. -// Instantiating the single-pass twin would only add a pipeline nothing can -// dispatch. +// launched at the Metal maximum 1024 threads per threadgroup (`group_dims` in +// `sdpa_vector`, scaled_dot_product_attention.cpp). At D = 512 that is 48 +// live floats per thread against 1024 threads, an occupancy claim the 2-pass +// kernel does not make (it runs 32 x gqa_factor threads with 32 live floats), +// and the host routes every D = 512 call to the 2-pass form for exactly that +// reason. Instantiating the single-pass twin would only add a pipeline +// nothing can dispatch. #define instantiate_sdpa_vector_2pass(type, qk_dim, value_dim) \ instantiate_kernel( \ - "sdpa_vector_2pass_1_" #type "_" #qk_dim "_" #value_dim, \ + "sdpa_vector_2pass_fp32partials_1_" #type "_" #qk_dim "_" #value_dim, \ sdpa_vector_2pass_1, \ type, \ qk_dim, \ diff --git a/python/tests/test_fast_sdpa.py b/python/tests/test_fast_sdpa.py index 3a9c04b075..6512e27683 100644 --- a/python/tests/test_fast_sdpa.py +++ b/python/tests/test_fast_sdpa.py @@ -1,5 +1,7 @@ import math import os +import subprocess +import sys import unittest from itertools import product from unittest.mock import patch @@ -357,6 +359,55 @@ def test_sdpa_vector_gqa_long(self): out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale) self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + def test_sdpa_vector_head_dim_512(self): + # Head dim 512 takes the 2-pass vector kernel at every key length. + D = 512 + scale = D**-0.5 + mx.random.seed(0) + for (Nq, Nkv), L, dtype in product( + [(16, 2), (4, 4)], [7, 1000, 8192], [mx.float32, mx.float16] + ): + with self.subTest(Nq=Nq, Nkv=Nkv, L=L, dtype=dtype): + q = 5e-1 * mx.random.normal(shape=(1, Nq, 1, D), dtype=dtype) + k = 5e-1 * mx.random.normal(shape=(1, Nkv, L, D), dtype=dtype) + v = 5e-1 * mx.random.normal(shape=(1, Nkv, L, D), dtype=dtype) + for m in [None, mx.random.uniform(shape=(Nq, 1, L)) > 0.2]: + ref = mlx_ref_attn(q, k, v, scale, mask=m) + out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, mask=m + ) + tol = 1e-4 if dtype == mx.float32 else 2e-3 + self.assertTrue(mx.allclose(ref, out, atol=tol, rtol=tol)) + + @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") + def test_sdpa_vector_head_dim_512_gqa_dedup(self): + # The process reads the dedup switch once, so the case runs in a + # child process with the switch on. + script = """ +import mlx.core as mx +from test_fast_sdpa import mlx_ref_attn + +D = 512 +mx.random.seed(0) +for L in [8192, 8201]: + for dtype, tol in [(mx.float32, 1e-4), (mx.float16, 2e-3)]: + q = 5e-1 * mx.random.normal(shape=(1, 16, 1, D), dtype=dtype) + k = 5e-1 * mx.random.normal(shape=(1, 2, L, D), dtype=dtype) + v = 5e-1 * mx.random.normal(shape=(1, 2, L, D), dtype=dtype) + ref = mlx_ref_attn(q, k, v, D**-0.5) + out = mx.fast.scaled_dot_product_attention(q, k, v, scale=D**-0.5) + assert mx.allclose(ref, out, atol=tol, rtol=tol), (L, dtype) +""" + env = dict(os.environ, DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP="1") + result = subprocess.run( + [sys.executable, "-c", script], + cwd=os.path.dirname(os.path.abspath(__file__)), + env=env, + capture_output=True, + text=True, + ) + self.assertEqual(result.returncode, 0, result.stderr) + @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") def test_sdpa_two_pass_partial_cancellation(self): # Uniform attention has the exact output 1 / L. Casting an From 9e239f9e8deb4af3c1720d85c347457e0fe96eea Mon Sep 17 00:00:00 2001 From: David Tai Date: Sun, 4 Oct 2026 20:55:25 -0500 Subject: [PATCH 08/11] docs(forkdiff): describe the Gemma 4 kernel work in fork.yaml 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 --- fork.yaml | 29 +++++++++++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/fork.yaml b/fork.yaml index 0bf5f5bc56..c2e22cbac1 100644 --- a/fork.yaml +++ b/fork.yaml @@ -119,6 +119,35 @@ def: - "mlx/backend/metal/kernels/scaled_dot_product_attention.metal" - "mlx/backend/metal/scaled_dot_product_attention.cpp" - "python/tests/test_fast_sdpa.py" + - title: "Gemma 4 decode and prefill kernels" + description: | + Kernel work for Gemma 4 26B-A4B, ported from the Gemma 4 engine repository. (PR #13) + - Softmax and vector SDPA: `#pragma unroll` on the loops with a fixed trip count. + The output does not change. + - Steel fused GEMM (`steel_gemm_fused.h`, `steel_gemm_fused_nax.h`): the addmm + epilogue identifies the composed-prefill causal-bias operand by its layout + (bf16, `ldc == N + 1`, `M <= N`, zero batch strides) and computes its two constant + values instead of loading them. The stored words do not change. + - `affine_gather_qmm_rhs_nax`: a simdgroup skips the A loads and MMAs for fragment + rows outside the current expert segment. + - `quantized.h`: affine 4-bit and 8-bit QMV kernels for the Gemma 4 decode shapes + (paired and multi-stream expert kernels, cross-row kernels, an 8x8 simdgroup-MMA + kernel for 8-row decode, and a tiled down-projection gather). + - Vector SDPA at head dim 512: the host sends every head-dim 512 vector call to the + 2-pass kernel. `DARKBLOOM_GEMMA4_D512_DECODE_2PASS=0` restores the unfused graph. + `DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP=1` (off by default) selects the GQA + kernel with 2 heads for each simdgroup. That kernel writes its merge plane in 2 + passes, so it stays inside the 32 KB threadgroup memory limit. + globs: + - "mlx/backend/metal/kernels/softmax.h" + - "mlx/backend/metal/kernels/sdpa_vector.h" + - "mlx/backend/metal/kernels/scaled_dot_product_attention.metal" + - "mlx/backend/metal/scaled_dot_product_attention.cpp" + - "mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h" + - "mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused_nax.h" + - "mlx/backend/metal/kernels/quantized_nax.h" + - "mlx/backend/metal/kernels/quantized.h" + - "python/tests/test_fast_sdpa.py" - title: "Declared-mutable inputs for Metal custom kernels" description: | `metal_kernel_with_mutable_inputs` lets a caller declare which custom-kernel inputs will From e3905de719c650cf47f2eb120b75dfbd9a73971c Mon Sep 17 00:00:00 2001 From: David Tai Date: Mon, 5 Oct 2026 00:15:48 -0500 Subject: [PATCH 09/11] refactor: shorten the D512 comment and tidy the head-dim 512 tests 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 --- .../metal/kernels/scaled_dot_product_attention.metal | 12 ++---------- python/tests/test_fast_sdpa.py | 7 ++++--- 2 files changed, 6 insertions(+), 13 deletions(-) diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal index 62c71c5935..a5f132507f 100644 --- a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal +++ b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal @@ -43,16 +43,8 @@ using namespace metal; hpt, \ split) -// D512-2PASS. The 2-pass split-K kernels ONLY, with no single-pass twin. -// -// `sdpa_vector` holds q, k and o in registers at D / 32 floats each and is -// launched at the Metal maximum 1024 threads per threadgroup (`group_dims` in -// `sdpa_vector`, scaled_dot_product_attention.cpp). At D = 512 that is 48 -// live floats per thread against 1024 threads, an occupancy claim the 2-pass -// kernel does not make (it runs 32 x gqa_factor threads with 32 live floats), -// and the host routes every D = 512 call to the 2-pass form for exactly that -// reason. Instantiating the single-pass twin would only add a pipeline -// nothing can dispatch. +// D512-2PASS. 2-pass kernels only. The host never runs `sdpa_vector` at +// D = 512, because at 1024 threads it would hold 48 floats per thread. #define instantiate_sdpa_vector_2pass(type, qk_dim, value_dim) \ instantiate_kernel( \ "sdpa_vector_2pass_fp32partials_1_" #type "_" #qk_dim "_" #value_dim, \ diff --git a/python/tests/test_fast_sdpa.py b/python/tests/test_fast_sdpa.py index 6512e27683..56afcd42a7 100644 --- a/python/tests/test_fast_sdpa.py +++ b/python/tests/test_fast_sdpa.py @@ -368,6 +368,7 @@ def test_sdpa_vector_head_dim_512(self): [(16, 2), (4, 4)], [7, 1000, 8192], [mx.float32, mx.float16] ): with self.subTest(Nq=Nq, Nkv=Nkv, L=L, dtype=dtype): + tol = 1e-4 if dtype == mx.float32 else 2e-3 q = 5e-1 * mx.random.normal(shape=(1, Nq, 1, D), dtype=dtype) k = 5e-1 * mx.random.normal(shape=(1, Nkv, L, D), dtype=dtype) v = 5e-1 * mx.random.normal(shape=(1, Nkv, L, D), dtype=dtype) @@ -376,7 +377,6 @@ def test_sdpa_vector_head_dim_512(self): out = mx.fast.scaled_dot_product_attention( q, k, v, scale=scale, mask=m ) - tol = 1e-4 if dtype == mx.float32 else 2e-3 self.assertTrue(mx.allclose(ref, out, atol=tol, rtol=tol)) @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") @@ -388,14 +388,15 @@ def test_sdpa_vector_head_dim_512_gqa_dedup(self): from test_fast_sdpa import mlx_ref_attn D = 512 +scale = D**-0.5 mx.random.seed(0) for L in [8192, 8201]: for dtype, tol in [(mx.float32, 1e-4), (mx.float16, 2e-3)]: q = 5e-1 * mx.random.normal(shape=(1, 16, 1, D), dtype=dtype) k = 5e-1 * mx.random.normal(shape=(1, 2, L, D), dtype=dtype) v = 5e-1 * mx.random.normal(shape=(1, 2, L, D), dtype=dtype) - ref = mlx_ref_attn(q, k, v, D**-0.5) - out = mx.fast.scaled_dot_product_attention(q, k, v, scale=D**-0.5) + ref = mlx_ref_attn(q, k, v, scale) + out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale) assert mx.allclose(ref, out, atol=tol, rtol=tol), (L, dtype) """ env = dict(os.environ, DARKBLOOM_GEMMA4_D512_DECODE_2PASS_DEDUP="1") From e602e6123928f816636025539b98a7e6251d5cdd Mon Sep 17 00:00:00 2001 From: David Tai Date: Mon, 5 Oct 2026 04:17:15 -0500 Subject: [PATCH 10/11] fix(metal): sum the QMV bias run in float, as load_vector does since #15 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 --- fork.yaml | 1 + mlx/backend/metal/kernels/quantized.h | 27 ++++++++++++++------------- python/tests/test_quantized.py | 15 +++++++++++++++ 3 files changed, 30 insertions(+), 13 deletions(-) diff --git a/fork.yaml b/fork.yaml index c2e22cbac1..5e3dafd4cb 100644 --- a/fork.yaml +++ b/fork.yaml @@ -148,6 +148,7 @@ def: - "mlx/backend/metal/kernels/quantized_nax.h" - "mlx/backend/metal/kernels/quantized.h" - "python/tests/test_fast_sdpa.py" + - "python/tests/test_quantized.py" - title: "Declared-mutable inputs for Metal custom kernels" description: | `metal_kernel_with_mutable_inputs` lets a caller declare which custom-kernel inputs will diff --git a/mlx/backend/metal/kernels/quantized.h b/mlx/backend/metal/kernels/quantized.h index bea5d39e36..ec0fd0b4dc 100644 --- a/mlx/backend/metal/kernels/quantized.h +++ b/mlx/backend/metal/kernels/quantized.h @@ -1272,10 +1272,11 @@ METAL_FUNC void qmv_fast_crossrow_affine4_g64_wide( // multiplied by the UNSCALED activation: (x / 4^k) * (w & (3 << 2k)) and // x * ((w >> 2k) & 3) are the same real product (power-of-two scaling is // exact in FP32), so every elementary product equals the generic -// qmv_fast_impl value; the wider lane coverage reassociates the -// FP32 partial sums, which is safe for this stage because the coarse -// shortlist is approximate by design and the exact affine-4 rerank plus -// target verification decide every emitted token. The serial leg runs no +// qmv_fast_impl value. The bias run sum widens each x to FP32 +// before the add, as load_vector does, so the wider lane coverage only +// reassociates FP32 partial sums. That is safe for this stage because the +// coarse shortlist is approximate by design and the exact affine-4 rerank +// plus target verification decide every emitted token. The serial leg runs no // 2-bit matmul (all its projections are affine-4), and out_vec_size == // 98_336 exists only in the compact draft readout, so the dispatch gate // below cannot touch the serial numerator or the denominator band. @@ -1329,7 +1330,8 @@ METAL_FUNC void qmv_fast_singlerow_affine2_g64( x0[i + 1] = static_cast(xm[i + 1]); x0[i + 2] = static_cast(xm[i + 2]); x0[i + 3] = static_cast(xm[i + 3]); - sum += xm[i] + xm[i + 1] + xm[i + 2] + xm[i + 3]; + sum += + float(xm[i]) + float(xm[i + 1]) + float(xm[i + 2]) + float(xm[i + 3]); } for (int r = 0; r < rows_per_simd; r++) { @@ -2279,10 +2281,9 @@ inline float mma8_hi(uint u) { } // Textual twin of `load_vector`'s `sum` on the same aligned -// 8-run that the reference lane owns: the parenthesised 4-tuple is evaluated -// on T exactly as in the reference, then the two trees are added in fp32. The -// bias term of the affine form therefore reuses the reference's own -// elementary values, not a re-derived sum. +// 8-run that the reference lane owns: each value widens to fp32 before the +// 4-tuple add, as in the reference. The bias term of the affine form +// therefore reuses the reference's own sum, not a re-derived one. template inline float mma8_runsum4(uint4 r) { thread T xt[8]; @@ -2295,8 +2296,8 @@ inline float mma8_runsum4(uint4 r) { xt[6] = mma8_u16::cast(ushort(r.w & 0xFFFFu)); xt[7] = mma8_u16::cast(ushort(r.w >> 16)); float sum = 0; - sum += xt[0] + xt[1] + xt[2] + xt[3]; - sum += xt[4] + xt[5] + xt[6] + xt[7]; + sum += float(xt[0]) + float(xt[1]) + float(xt[2]) + float(xt[3]); + sum += float(xt[4]) + float(xt[5]) + float(xt[6]) + float(xt[7]); return sum; } @@ -3383,8 +3384,8 @@ template < // a g64 group and sums them before the single `s * C + rs * b` close. Every // elementary term is the reference's own -- the products x * q are exact in // fp32 (a bf16 x carries 8 significant bits, a code 4), scales and biases - // widen exactly, `mma8_runsum4` reproduces `load_vector`'s bf16 4-tuple sum - // order on the same aligned 8-run, and the group closes are chained in + // widen exactly, `mma8_runsum4` reproduces `load_vector`'s fp32 4-tuple sum + // on the same aligned 8-run, and the group closes are chained in // ascending k -- so the ONLY numeric deviation is fp32 reassociation inside // the 64-wide group dot (plus the two-halves add of the KS = 2 split). This // is the first non-bit-exact QMV tier here; measured against the stock M = 1 diff --git a/python/tests/test_quantized.py b/python/tests/test_quantized.py index 461175f013..620f1597cd 100644 --- a/python/tests/test_quantized.py +++ b/python/tests/test_quantized.py @@ -710,6 +710,21 @@ def test_fp_qmv_large_output(self): tol = 1e-2 if dtype == mx.bfloat16 else 1e-3 self.assertTrue(mx.allclose(y_q, y_hat, rtol=tol, atol=tol)) + def test_qmv_bias_sum_widens_inputs(self): + if mx.default_device() == mx.cpu: + self.skipTest("Checks the Metal QMV tiers") + # In bf16, 256 + 1 rounds to 256, so a bias sum of the 4-tuple + # (256, 1, 1, 1) taken in T loses 3 per tuple. Every tier must add in fp32. + for M, N, K, bits in [(1, 98336, 1024, 2), (8, 1024, 2816, 4)]: + with self.subTest(M=M, N=N, K=K, bits=bits): + x = mx.tile(mx.array([256, 1, 1, 1], mx.bfloat16), (M, K // 4)) + w = mx.zeros((N, K * bits // 32), mx.uint32) + scales = mx.ones((N, K // 64), mx.bfloat16) + biases = mx.ones((N, K // 64), mx.bfloat16) + y = mx.quantized_matmul(x, w, scales, biases, True, 64, bits) + expected = mx.full((M, N), 259 * K // 4, mx.float32) + self.assertTrue(mx.array_equal(y, expected.astype(mx.bfloat16))) + def test_qmv_wide(self): # M in [2, vector_limit) routes to qmv_wide -- except K in {64, 128} # with power-of-2 bits, which stays on qmv_quad. Check both paths From 725859207e4bc0db16e447b198e1c171574321e4 Mon Sep 17 00:00:00 2001 From: David Tai Date: Tue, 6 Oct 2026 15:18:21 -0500 Subject: [PATCH 11/11] style: clang-format the Metal kernel and SDPA files (pre-commit v21.1.8) Formatting only. Stripping comments and whitespace gives identical code for all six files. Co-Authored-By: Claude Opus 5.5 --- mlx/backend/metal/kernels/quantized.h | 519 ++++++++++++------ mlx/backend/metal/kernels/quantized_nax.h | 28 +- mlx/backend/metal/kernels/sdpa_vector.h | 44 +- mlx/backend/metal/kernels/softmax.h | 24 +- .../steel/gemm/kernels/steel_gemm_fused.h | 63 +-- .../metal/scaled_dot_product_attention.cpp | 3 +- 6 files changed, 412 insertions(+), 269 deletions(-) diff --git a/mlx/backend/metal/kernels/quantized.h b/mlx/backend/metal/kernels/quantized.h index ec0fd0b4dc..2dd41bc90e 100644 --- a/mlx/backend/metal/kernels/quantized.h +++ b/mlx/backend/metal/kernels/quantized.h @@ -327,15 +327,11 @@ inline U qdot_affine4_registered_word( const uint packed0 = packed_word & 0xffffu; const uint packed1 = packed_word >> 16; U accum = - (x_thread[0] * (packed0 & 0x000f) + - x_thread[1] * (packed0 & 0x00f0) + - x_thread[2] * (packed0 & 0x0f00) + - x_thread[3] * (packed0 & 0xf000)); + (x_thread[0] * (packed0 & 0x000f) + x_thread[1] * (packed0 & 0x00f0) + + x_thread[2] * (packed0 & 0x0f00) + x_thread[3] * (packed0 & 0xf000)); accum += - (x_thread[4] * (packed1 & 0x000f) + - x_thread[5] * (packed1 & 0x00f0) + - x_thread[6] * (packed1 & 0x0f00) + - x_thread[7] * (packed1 & 0xf000)); + (x_thread[4] * (packed1 & 0x000f) + x_thread[5] * (packed1 & 0x00f0) + + x_thread[6] * (packed1 & 0x0f00) + x_thread[7] * (packed1 & 0xf000)); return scale * accum + sum * bias; } @@ -358,30 +354,23 @@ inline void qdot_affine4_pair( const uint packed0 = packedWord & 0xffffu; const uint packed1 = packedWord >> 16; U accum0 = - (x0[0] * (packed0 & 0x000f) + - x0[1] * (packed0 & 0x00f0) + - x0[2] * (packed0 & 0x0f00) + - x0[3] * (packed0 & 0xf000)); + (x0[0] * (packed0 & 0x000f) + x0[1] * (packed0 & 0x00f0) + + x0[2] * (packed0 & 0x0f00) + x0[3] * (packed0 & 0xf000)); U accum1 = - (x1[0] * (packed0 & 0x000f) + - x1[1] * (packed0 & 0x00f0) + - x1[2] * (packed0 & 0x0f00) + - x1[3] * (packed0 & 0xf000)); + (x1[0] * (packed0 & 0x000f) + x1[1] * (packed0 & 0x00f0) + + x1[2] * (packed0 & 0x0f00) + x1[3] * (packed0 & 0xf000)); accum0 += - (x0[4] * (packed1 & 0x000f) + - x0[5] * (packed1 & 0x00f0) + - x0[6] * (packed1 & 0x0f00) + - x0[7] * (packed1 & 0xf000)); + (x0[4] * (packed1 & 0x000f) + x0[5] * (packed1 & 0x00f0) + + x0[6] * (packed1 & 0x0f00) + x0[7] * (packed1 & 0xf000)); accum1 += - (x1[4] * (packed1 & 0x000f) + - x1[5] * (packed1 & 0x00f0) + - x1[6] * (packed1 & 0x0f00) + - x1[7] * (packed1 & 0xf000)); + (x1[4] * (packed1 & 0x000f) + x1[5] * (packed1 & 0x00f0) + + x1[6] * (packed1 & 0x0f00) + x1[7] * (packed1 & 0xf000)); out0 = scale * accum0 + sum0 * bias; out1 = scale * accum1 + sum1 * bias; } -// Two independent affine-4 dot products over one register-held packed 32-bit word. +// Two independent affine-4 dot products over one register-held packed 32-bit +// word. template inline void qdot_affine4_pair_word( uint packedWord, @@ -397,25 +386,17 @@ inline void qdot_affine4_pair_word( const uint packed0 = packedWord & 0xffffu; const uint packed1 = packedWord >> 16; U accum0 = - (x0[0] * (packed0 & 0x000f) + - x0[1] * (packed0 & 0x00f0) + - x0[2] * (packed0 & 0x0f00) + - x0[3] * (packed0 & 0xf000)); + (x0[0] * (packed0 & 0x000f) + x0[1] * (packed0 & 0x00f0) + + x0[2] * (packed0 & 0x0f00) + x0[3] * (packed0 & 0xf000)); U accum1 = - (x1[0] * (packed0 & 0x000f) + - x1[1] * (packed0 & 0x00f0) + - x1[2] * (packed0 & 0x0f00) + - x1[3] * (packed0 & 0xf000)); + (x1[0] * (packed0 & 0x000f) + x1[1] * (packed0 & 0x00f0) + + x1[2] * (packed0 & 0x0f00) + x1[3] * (packed0 & 0xf000)); accum0 += - (x0[4] * (packed1 & 0x000f) + - x0[5] * (packed1 & 0x00f0) + - x0[6] * (packed1 & 0x0f00) + - x0[7] * (packed1 & 0xf000)); + (x0[4] * (packed1 & 0x000f) + x0[5] * (packed1 & 0x00f0) + + x0[6] * (packed1 & 0x0f00) + x0[7] * (packed1 & 0xf000)); accum1 += - (x1[4] * (packed1 & 0x000f) + - x1[5] * (packed1 & 0x00f0) + - x1[6] * (packed1 & 0x0f00) + - x1[7] * (packed1 & 0xf000)); + (x1[4] * (packed1 & 0x000f) + x1[5] * (packed1 & 0x00f0) + + x1[6] * (packed1 & 0x0f00) + x1[7] * (packed1 & 0xf000)); out0 = scale * accum0 + sum0 * bias; out1 = scale * accum1 + sum1 * bias; } @@ -1098,16 +1079,13 @@ METAL_FUNC void qmv_fast_crossrow_affine4_g64( for (int r = 0; r < rows_per_simd; r++) { const int row = out_row + r; - const device uint8_t* wb = - reinterpret_cast(w) + + const device uint8_t* wb = reinterpret_cast(w) + row * in_vec_size_w + k / 2 + simd_lid * bytes_per_lane; - const device uint16_t* ws = - reinterpret_cast(wb); + const device uint16_t* ws = reinterpret_cast(wb); for (int i = 0; i < 4; i++) { packed[r][i] = ws[i]; } - const int group_index = - row * in_vec_size_g + k / 64 + simd_lid / 4; + const int group_index = row * in_vec_size_g + k / 64 + simd_lid / 4; scale_local[r] = scales[group_index]; bias_local[r] = biases[group_index]; } @@ -1115,16 +1093,18 @@ METAL_FUNC void qmv_fast_crossrow_affine4_g64( thread float x0[values_per_thread]; const device T* xm0 = x + first_m * in_vec_size + k + simd_lid * values_per_thread; - const float sum0 = - load_vector(xm0, x0); + const float sum0 = load_vector(xm0, x0); if (has_pair) { thread float x1[values_per_thread]; const device T* xm1 = xm0 + in_vec_size; - const float sum1 = - load_vector(xm1, x1); + const float sum1 = load_vector(xm1, x1); for (int r = 0; r < rows_per_simd; r++) { pair_result[r] += qdot_affine4_loaded_pair( - packed[r], x0, x1, scale_local[r], bias_local[r], + packed[r], + x0, + x1, + scale_local[r], + bias_local[r], float2(sum0, sum1)); } } else { @@ -1236,15 +1216,15 @@ METAL_FUNC void qmv_fast_crossrow_affine4_g64_wide( } for (int r = 0; r < rows_per_simd; r++) { if (DIRECT_NIBBLES) { - partial[r] += (a0 * (packed[r][i] & 0x000f) + - a1 * ((packed[r][i] >> 4) & 0x000f) + - a2 * ((packed[r][i] >> 8) & 0x000f) + - a3 * ((packed[r][i] >> 12) & 0x000f)); + partial[r] += + (a0 * (packed[r][i] & 0x000f) + + a1 * ((packed[r][i] >> 4) & 0x000f) + + a2 * ((packed[r][i] >> 8) & 0x000f) + + a3 * ((packed[r][i] >> 12) & 0x000f)); } else { - partial[r] += (a0 * (packed[r][i] & 0x000f) + - a1 * (packed[r][i] & 0x00f0) + - a2 * (packed[r][i] & 0x0f00) + - a3 * (packed[r][i] & 0xf000)); + partial[r] += + (a0 * (packed[r][i] & 0x000f) + a1 * (packed[r][i] & 0x00f0) + + a2 * (packed[r][i] & 0x0f00) + a3 * (packed[r][i] & 0xf000)); } } } @@ -1257,8 +1237,7 @@ METAL_FUNC void qmv_fast_crossrow_affine4_g64_wide( for (int m = 0; m < NA; m++) { const float reduced = simd_sum(acc[r][m]); if (simd_lid == 0) { - y[(first_m + m) * out_vec_size + out_row + r] = - static_cast(reduced); + y[(first_m + m) * out_vec_size + out_row + r] = static_cast(reduced); } } } @@ -1295,9 +1274,9 @@ METAL_FUNC void qmv_fast_singlerow_affine2_g64( constexpr int rows_per_simd = 4; constexpr int values_per_thread = 32; constexpr int block_size = values_per_thread * SIMD_SIZE; - constexpr int bytes_per_lane = 8; // 32 values x 2 bits = 8 bytes - const int in_vec_size_w = in_vec_size / 4; // weight bytes per output row - const int in_vec_size_g = in_vec_size / 64; // scale groups per output row + constexpr int bytes_per_lane = 8; // 32 values x 2 bits = 8 bytes + const int in_vec_size_w = in_vec_size / 4; // weight bytes per output row + const int in_vec_size_g = in_vec_size / 64; // scale groups per output row const int out_row = int(tid.y) * 8 + int(simd_gid) * rows_per_simd; @@ -1336,7 +1315,7 @@ METAL_FUNC void qmv_fast_singlerow_affine2_g64( for (int r = 0; r < rows_per_simd; r++) { float accum = 0.0f; - #pragma unroll +#pragma unroll for (int j = 0; j < 32; j++) { accum += x0[j] * float((packed[r] >> (2 * j)) & 0x03ul); } @@ -1366,7 +1345,8 @@ METAL_FUNC void qmv_fast_crossrow_affine4_g64_m( uint3 tid, uint simd_gid, uint simd_lid) { - static_assert(M >= 3 && M <= 9, "wide multi-row QMV dispatch covers M in [3, 9]"); + static_assert( + M >= 3 && M <= 9, "wide multi-row QMV dispatch covers M in [3, 9]"); static_assert(M % IPG != 1, "a one-input tail group is not instantiated"); constexpr int TAIL = M % IPG; const int first_m = int(tid.x) * IPG; @@ -1376,13 +1356,31 @@ METAL_FUNC void qmv_fast_crossrow_affine4_g64_m( const int out_row = int(tid.y) * 8 + int(simd_gid) * 4; if (TAIL == 0 || M - first_m >= IPG) { qmv_fast_crossrow_affine4_g64_wide( - w, scales, biases, x, y, in_vec_size, out_vec_size, - first_m, out_row, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + first_m, + out_row, + simd_lid); } else { qmv_fast_crossrow_affine4_g64_wide< - T, (TAIL >= 2 ? TAIL : 2), DIRECT_NIBBLES>( - w, scales, biases, x, y, in_vec_size, out_vec_size, - first_m, out_row, simd_lid); + T, + (TAIL >= 2 ? TAIL : 2), + DIRECT_NIBBLES>( + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + first_m, + out_row, + simd_lid); } } @@ -1632,7 +1630,15 @@ METAL_FUNC void qmv_affine4_g64_pair_impl( float dot0; float dot1; qdot_affine4_pair_word( - packed[row], x0_thread, x1_thread, scale_local[row], bias_local[row], sum0, sum1, dot0, dot1); + packed[row], + x0_thread, + x1_thread, + scale_local[row], + bias_local[row], + sum0, + sum1, + dot0, + dot1); result0[row] += dot0; result1[row] += dot1; } @@ -1648,8 +1654,7 @@ METAL_FUNC void qmv_affine4_g64_pair_impl( // 64. The final block therefore contains an integral number of complete // eight-value lane packets (32 lanes for K=2816, 24 for expert down_proj // K=704); no active lane needs the generic dynamic safe-tail loops. - const uint active_tail_lanes = - uint((in_vec_size - k) / values_per_thread); + const uint active_tail_lanes = uint((in_vec_size - k) / values_per_thread); if (simd_lid < active_tail_lanes) { for (int row = 0; row < results_per_simdgroup; row++) { packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); @@ -1657,15 +1662,21 @@ METAL_FUNC void qmv_affine4_g64_pair_impl( bias_local[row] = biases[row * in_vec_size_g]; } - float sum0 = - load_vector(x0, x0_thread); - float sum1 = - load_vector(x1, x1_thread); + float sum0 = load_vector(x0, x0_thread); + float sum1 = load_vector(x1, x1_thread); for (int row = 0; row < results_per_simdgroup; row++) { float dot0; float dot1; qdot_affine4_pair_word( - packed[row], x0_thread, x1_thread, scale_local[row], bias_local[row], sum0, sum1, dot0, dot1); + packed[row], + x0_thread, + x1_thread, + scale_local[row], + bias_local[row], + sum0, + sum1, + dot0, + dot1); result0[row] += dot0; result1[row] += dot1; } @@ -1758,8 +1769,7 @@ METAL_FUNC void qmv_affine4_g64_quad_stream_impl( int k = 0; for (; k <= in_vec_size - block_size; k += block_size) { for (int row = 0; row < results_per_simdgroup; row++) { - packed[row] = - *((const device uint*)(ws + row * in_vec_size_w)); + packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); scale_local[row] = scales[row * in_vec_size_g]; bias_local[row] = biases[row * in_vec_size_g]; } @@ -1800,32 +1810,31 @@ METAL_FUNC void qmv_affine4_g64_quad_stream_impl( values_per_thread); if (remaining > 0) { for (int row = 0; row < results_per_simdgroup; row++) { - packed[row] = - *((const device uint*)(ws + row * in_vec_size_w)); + packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); scale_local[row] = scales[row * in_vec_size_g]; bias_local[row] = biases[row * in_vec_size_g]; } - float sum = - load_vector_safe(x0, x_thread, remaining); + float sum = load_vector_safe( + x0, x_thread, remaining); for (int row = 0; row < results_per_simdgroup; row++) { result0[row] += qdot_affine4_registered_word( packed[row], x_thread, scale_local[row], bias_local[row], sum); } - sum = - load_vector_safe(x1, x_thread, remaining); + sum = load_vector_safe( + x1, x_thread, remaining); for (int row = 0; row < results_per_simdgroup; row++) { result1[row] += qdot_affine4_registered_word( packed[row], x_thread, scale_local[row], bias_local[row], sum); } - sum = - load_vector_safe(x2, x_thread, remaining); + sum = load_vector_safe( + x2, x_thread, remaining); for (int row = 0; row < results_per_simdgroup; row++) { result2[row] += qdot_affine4_registered_word( packed[row], x_thread, scale_local[row], bias_local[row], sum); } - sum = - load_vector_safe(x3, x_thread, remaining); + sum = load_vector_safe( + x3, x_thread, remaining); for (int row = 0; row < results_per_simdgroup; row++) { result3[row] += qdot_affine4_registered_word( packed[row], x_thread, scale_local[row], bias_local[row], sum); @@ -1900,8 +1909,7 @@ METAL_FUNC void qmv_affine4_g64_triple_stream_impl( int k = 0; for (; k <= in_vec_size - block_size; k += block_size) { for (int row = 0; row < results_per_simdgroup; row++) { - packed[row] = - *((const device uint*)(ws + row * in_vec_size_w)); + packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); scale_local[row] = scales[row * in_vec_size_g]; bias_local[row] = biases[row * in_vec_size_g]; } @@ -1936,26 +1944,25 @@ METAL_FUNC void qmv_affine4_g64_triple_stream_impl( values_per_thread); if (remaining > 0) { for (int row = 0; row < results_per_simdgroup; row++) { - packed[row] = - *((const device uint*)(ws + row * in_vec_size_w)); + packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); scale_local[row] = scales[row * in_vec_size_g]; bias_local[row] = biases[row * in_vec_size_g]; } - float sum = - load_vector_safe(x0, x_thread, remaining); + float sum = load_vector_safe( + x0, x_thread, remaining); for (int row = 0; row < results_per_simdgroup; row++) { result0[row] += qdot_affine4_registered_word( packed[row], x_thread, scale_local[row], bias_local[row], sum); } - sum = - load_vector_safe(x1, x_thread, remaining); + sum = load_vector_safe( + x1, x_thread, remaining); for (int row = 0; row < results_per_simdgroup; row++) { result1[row] += qdot_affine4_registered_word( packed[row], x_thread, scale_local[row], bias_local[row], sum); } - sum = - load_vector_safe(x2, x_thread, remaining); + sum = load_vector_safe( + x2, x_thread, remaining); for (int row = 0; row < results_per_simdgroup; row++) { result2[row] += qdot_affine4_registered_word( packed[row], x_thread, scale_local[row], bias_local[row], sum); @@ -2178,8 +2185,7 @@ METAL_FUNC void qmv_affine8_g64_quad_stream_impl( // four-value lane packets. In particular down_proj K=2112 leaves exactly // 16 active lanes; use the fixed unrolled load instead of four dynamic // safe-tail loops while preserving each lane's qdot and simd_sum order. - const uint active_tail_lanes = - uint((in_vec_size - k) / values_per_thread); + const uint active_tail_lanes = uint((in_vec_size - k) / values_per_thread); if (simd_lid < active_tail_lanes) { for (int row = 0; row < results_per_simdgroup; row++) { packed[row] = *((const device uint*)(ws + row * in_vec_size_w)); @@ -2301,8 +2307,8 @@ inline float mma8_runsum4(uint4 r) { return sum; } -#define MMA8_SETB(BB, W, HI) \ - BB.thread_elements()[0] = mma8_##HI(r0.W); \ +#define MMA8_SETB(BB, W, HI) \ + BB.thread_elements()[0] = mma8_##HI(r0.W); \ BB.thread_elements()[1] = mma8_##HI(r1.W); #define MMA8_STEP(BB, J) \ @@ -3198,7 +3204,15 @@ template < // M == 1 coarse draft readout (draft-rerank scheme): the ONE 2-bit shape // in the scored path; proposal-only by construction (see kernel header). qmv_fast_singlerow_affine2_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, tid, simd_gid, + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, simd_lid); return; } @@ -3210,56 +3224,120 @@ template < switch (ntg.x) { case 2: qmv_fast_crossrow_affine4_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 3: qmv_fast_crossrow_affine4_g64_m( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 4: qmv_fast_crossrow_affine4_g64_m( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 5: qmv_fast_crossrow_affine4_g64_m( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 6: qmv_fast_crossrow_affine4_g64_m( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 7: qmv_fast_crossrow_affine4_g64_m( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 8: // 3+3+2, not 4+4. M = 8 is the only hot width whose EVEN split needs // two simultaneous vec accumulators in every active worker; - // M = 9 uses three-lane vectors and profiles CHEAPER despite more work - // (319 / 437 / 216 us for M = 7 / 8 / 9 in the public cross-row study) - // — a register cliff, not work scaling. - // Exact: these lanes carry INDEPENDENT input rows and are never reduced - // across (simd_sum reduces along K WITHIN a row), so moving a row from - // lane 3 of a four-wide vector to lane 0 of a two-wide one cannot - // reorder its scalar chain. Template admits it: M in [3,9], 8 % 3 == 2 - // (no one-row tail), IPG 3 inside the wide helper's [2,4]. - // Receipts: 85d5bca3 2.91143, yzxoi 2.92675. - // SYNERGY with the streak gate above, which is why they ship together: - // gate 2 reaches the width-8 verify SOONER, so this kernel fires MORE. + // M = 9 uses three-lane vectors and profiles CHEAPER despite more + // work (319 / 437 / 216 us for M = 7 / 8 / 9 in the public cross-row + // study) — a register cliff, not work scaling. Exact: these lanes + // carry INDEPENDENT input rows and are never reduced across (simd_sum + // reduces along K WITHIN a row), so moving a row from lane 3 of a + // four-wide vector to lane 0 of a two-wide one cannot reorder its + // scalar chain. Template admits it: M in [3,9], 8 % 3 == 2 (no + // one-row tail), IPG 3 inside the wide helper's [2,4]. Receipts: + // 85d5bca3 2.91143, yzxoi 2.92675. SYNERGY with the streak gate + // above, which is why they ship together: gate 2 reaches the width-8 + // verify SOONER, so this kernel fires MORE. qmv_fast_crossrow_affine4_g64_m( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 9: qmv_fast_crossrow_affine4_g64_m( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; default: break; @@ -3268,43 +3346,107 @@ template < switch (ntg.x) { case 2: qmv_fast_crossrow_affine4_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 3: qmv_fast_crossrow_affine4_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 4: qmv_fast_crossrow_affine4_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 5: qmv_fast_crossrow_affine4_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 6: qmv_fast_crossrow_affine4_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 7: qmv_fast_crossrow_affine4_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 8: qmv_fast_crossrow_affine4_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; case 9: qmv_fast_crossrow_affine4_g64( - w, scales, biases, x, y, in_vec_size, out_vec_size, - tid, simd_gid, simd_lid); + w, + scales, + biases, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); return; default: break; @@ -3370,9 +3512,8 @@ template < b_strides, tid); } - if (!batched && group_size == 64 && bits == 4 && ntg.x == 8 && - ntg.z == 1 && in_vec_size % 64 == 0 && out_vec_size >= 8 && - out_vec_size % 8 == 0) { + if (!batched && group_size == 64 && bits == 4 && ntg.x == 8 && ntg.z == 1 && + in_vec_size % 64 == 0 && out_vec_size >= 8 && out_vec_size % 8 == 0) { // The ruled decode cohort presents eight input rows to ordinary QMV. // MMA-QKV S1 -- GROUP-EXACT-MMA tier. It replaces, for the 4-bit affine // g64 dense decode projections wide enough to fill the machine (q/k/v: @@ -3388,24 +3529,25 @@ template < // on the same aligned 8-run, and the group closes are chained in // ascending k -- so the ONLY numeric deviation is fp32 reassociation inside // the 64-wide group dot (plus the two-halves add of the KS = 2 split). This - // is the first non-bit-exact QMV tier here; measured against the stock M = 1 - // road over 50 random cohorts per plane at K = 2816, the deviation is at + // is the first non-bit-exact QMV tier here; measured against the stock M = + // 1 road over 50 random cohorts per plane at K = 2816, the deviation is at // most 1 bf16 ulp for every output above the 2^-10 * row-max magnitude gate // (non-zero fraction ~1.4e-4, 0 argmax flips over 400 rows per plane), it // is run-to-run bitwise deterministic, and the body measured 0.41-0.51x the - // quad_stream body net of the dispatch floor on an M4 Max. Outputs cancelled - // below ~2^-8 of their term mass can show a second relative ulp; they are - // numerically negligible and never argmax candidates. KILL SWITCH: set - // `kGemma4QmvMma8Affine4` to false and this branch vanishes at compile time, - // restoring the quad_stream and pair tiers below byte for byte -- nothing - // beneath this block was edited. Raising `kGemma4QmvMma8Affine4FloorN` - // returns individual planes the same way. (MSL forbids a program-scope - // `constexpr`, so the two switches live at the top of the tier they guard.) + // quad_stream body net of the dispatch floor on an M4 Max. Outputs + // cancelled below ~2^-8 of their term mass can show a second relative ulp; + // they are numerically negligible and never argmax candidates. KILL SWITCH: + // set `kGemma4QmvMma8Affine4` to false and this branch vanishes at compile + // time, restoring the quad_stream and pair tiers below byte for byte -- + // nothing beneath this block was edited. Raising + // `kGemma4QmvMma8Affine4FloorN` returns individual planes the same way. + // (MSL forbids a program-scope `constexpr`, so the two switches live at the + // top of the tier they guard.) constexpr bool kGemma4QmvMma8Affine4 = true; constexpr int kGemma4QmvMma8Affine4FloorN = 1024; if (kGemma4QmvMma8Affine4 && sizeof(T) == 2 && ntg.z == 1 && - in_vec_size % 64 == 0 && - out_vec_size >= kGemma4QmvMma8Affine4FloorN && out_vec_size % 8 == 0) { + in_vec_size % 64 == 0 && out_vec_size >= kGemma4QmvMma8Affine4FloorN && + out_vec_size % 8 == 0) { // Seven of the eight host x-groups retire before any load; the eighth // produces all eight cohort columns of its eight output rows. if (tid.x != 0) { @@ -3491,9 +3633,8 @@ template < simd_lid); return; } - if (!batched && group_size == 64 && bits == 8 && ntg.x == 8 && - ntg.z == 1 && in_vec_size % 64 == 0 && out_vec_size >= 8 && - out_vec_size % 8 == 0) { + if (!batched && group_size == 64 && bits == 8 && ntg.x == 8 && ntg.z == 1 && + in_vec_size % 64 == 0 && out_vec_size >= 8 && out_vec_size % 8 == 0) { // Dense decode projections use byte weights. if (out_vec_size >= 1024) { // WIDE-N tier -- the dense MLP of all 30 layers: gate_proj and up_proj @@ -4264,8 +4405,7 @@ METAL_FUNC void gather_qmv_gemma4_down_tile( const device uint32_t* tile_w = w + expert * w_stride; const device T* tile_scales = scales + expert * s_stride; const device T* tile_biases = biases + expert * b_stride; - const device T* tile_x0 = - x + lhs_indices[assignment * lhs_stride] * x_stride; + const device T* tile_x0 = x + lhs_indices[assignment * lhs_stride] * x_stride; device T* tile_y0 = y + assignment * out_vec_size; const bool has_pair = expert_prefix_bounds ? (((route_word >> 14) & 0x3fu) + 1u) > 1u @@ -4337,9 +4477,9 @@ template uint simd_gid [[simdgroup_index_in_threadgroup]], uint simd_lid [[thread_index_in_simdgroup]]) { int M = x_shape[x_batch_ndims]; - const bool gemma4_pair_geometry = - group_size == 64 && bits == 4 && M == 1 && batch_ndims == 1 && - batch_shape[0] == 64 && x_batch_ndims == 1 && w_batch_ndims == 1 && + const bool gemma4_pair_geometry = group_size == 64 && bits == 4 && M == 1 && + batch_ndims == 1 && batch_shape[0] == 64 && x_batch_ndims == 1 && + w_batch_ndims == 1 && ((in_vec_size == 2816 && out_vec_size == 704) || (in_vec_size == 704 && out_vec_size == 2816)); if (gemma4_pair_geometry) { @@ -4371,8 +4511,7 @@ template return; } const uint assignment = tid.z; - const uint32_t route_word = - rhs_indices[assignment * (uint)rhs_strides[0]]; + const uint32_t route_word = rhs_indices[assignment * (uint)rhs_strides[0]]; const bool expert_prefix_bounds = (route_word & 0x80000000u) != 0u; const uint32_t expert = expert_prefix_bounds ? (route_word & 0xffu) : route_word; @@ -4477,22 +4616,36 @@ template // Singleton experts and odd-run tails still need one ordinary QMV, but // their one-dimensional offsets are already resolved by this guard. - const uint32_t single_lhs = - lhs_indices[assignment * (uint)lhs_strides[0]]; + const uint32_t single_lhs = lhs_indices[assignment * (uint)lhs_strides[0]]; const device T* single_x = x + single_lhs * x_strides[0]; const device uint32_t* single_w = w + expert * w_strides[0]; const device T* single_scales = scales + expert * s_strides[0]; const device T* single_biases = biases + expert * b_strides[0]; device T* single_y = y + assignment * (uint)out_vec_size; if (in_vec_size == 2816) { - qmv_affine4_g64_singles_impl< - T, group_size, bits, 2816, true, false>( - single_w, single_scales, single_biases, single_x, single_y, - in_vec_size, out_vec_size, tid, simd_gid, simd_lid); + qmv_affine4_g64_singles_impl( + single_w, + single_scales, + single_biases, + single_x, + single_y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); } else { qmv_impl( - single_w, single_scales, single_biases, single_x, single_y, - in_vec_size, out_vec_size, tid, simd_gid, simd_lid); + single_w, + single_scales, + single_biases, + single_x, + single_y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); } return; } diff --git a/mlx/backend/metal/kernels/quantized_nax.h b/mlx/backend/metal/kernels/quantized_nax.h index 4a31c7e800..4aa5fdf254 100644 --- a/mlx/backend/metal/kernels/quantized_nax.h +++ b/mlx/backend/metal/kernels/quantized_nax.h @@ -1753,8 +1753,7 @@ template < // gather_qmm_rhs is dispatched only for right-sorted indices. If this // segment's expert matches the tile endpoint, sortedness proves that the // remaining suffix is one segment and the per-row probe can stop here. - if (kGatherRhsSortedEndpointElide && - indices[y_row + tgp_bm - 1] == index) { + if (kGatherRhsSortedEndpointElide && indices[y_row + tgp_bm - 1] == index) { n = tgp_bm; } else { for (; n < tgp_bm; n++) { @@ -1823,17 +1822,15 @@ template < volatile int compiler_barrier; if constexpr (transpose) { - Btile.template load( - Ws + tn * BK_padded + kk1); + Btile.template load(Ws + tn * BK_padded + kk1); } else { - Btile.template load( - Ws + tn + kk1 * BN_padded); + Btile.template load(Ws + tn + kk1 * BN_padded); } STEEL_PRAGMA_UNROLL for (short mm = 0; mm < TM; mm++) { const short fr = short(mm * Dtile.kFragRows); - if (fr < seg_hi && short(fr + Dtile.kFragRows) > seg_lo) { + if (fr seg_lo) { gather_rhs_load_frag_row(mm, Atile, xn + kk1, K); gather_rhs_mma_frag_row( mm, @@ -1861,11 +1858,9 @@ template < } if constexpr (transpose) { - Btile.template load( - Ws + tn * BK_padded + kk1); + Btile.template load(Ws + tn * BK_padded + kk1); } else { - Btile.template load( - Ws + tn + kk1 * BN_padded); + Btile.template load(Ws + tn + kk1 * BN_padded); } tile_matmad_nax( @@ -1903,11 +1898,9 @@ template < Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); if constexpr (transpose) { - Btile.template load( - Ws + tn * BK_padded + kk1); + Btile.template load(Ws + tn * BK_padded + kk1); } else { - Btile.template load( - Ws + tn + kk1 * BN_padded); + Btile.template load(Ws + tn + kk1 * BN_padded); } tile_matmad_nax( @@ -1936,10 +1929,7 @@ template < } } else { Dtile.store_slice( - y + tm * N + tn, - N, - short2(0, seg_lo), - short2(sgp_sn, seg_hi)); + y + tm * N + tn, N, short2(0, seg_lo), short2(sgp_sn, seg_hi)); } } }); diff --git a/mlx/backend/metal/kernels/sdpa_vector.h b/mlx/backend/metal/kernels/sdpa_vector.h index bb634eab94..6c73ea5d8a 100644 --- a/mlx/backend/metal/kernels/sdpa_vector.h +++ b/mlx/backend/metal/kernels/sdpa_vector.h @@ -80,12 +80,12 @@ template out += o_offset * V + simd_gid * v_per_thread; - // Read the query and 0 the output accumulator - #pragma unroll +// Read the query and 0 the output accumulator +#pragma unroll for (int i = 0; i < qk_per_thread; i++) { q[i] = static_cast(scale) * queries[i]; } - #pragma unroll +#pragma unroll for (int i = 0; i < v_per_thread; i++) { o[i] = 0; } @@ -108,15 +108,15 @@ template use_key = (fmask[0] >= Limits::finite_min); } if (use_key) { - // Read the key - #pragma unroll +// Read the key +#pragma unroll for (int j = 0; j < qk_per_thread; j++) { k[j] = keys[j]; } // Compute the i-th score U score = 0; - #pragma unroll +#pragma unroll for (int j = 0; j < qk_per_thread; j++) { score += q[j] * k[j]; } @@ -133,8 +133,8 @@ template max_score = new_max; sum_exp_score = sum_exp_score * factor + exp_score; - // Update the output accumulator - #pragma unroll +// Update the output accumulator +#pragma unroll for (int j = 0; j < v_per_thread; j++) { o[j] = o[j] * factor + exp_score * values[j]; } @@ -164,8 +164,8 @@ template U factor = fast::exp(max_score - new_max); sum_exp_score = simd_sum(sum_exp_scores[simd_lid] * factor); - // Now we need to aggregate all the outputs - #pragma unroll +// Now we need to aggregate all the outputs +#pragma unroll for (int i = 0; i < v_per_thread; i++) { outputs[simd_lid * BD + simd_gid] = o[i]; threadgroup_barrier(mem_flags::mem_threadgroup); @@ -176,7 +176,7 @@ template // And write the output if (simd_lid == 0) { - #pragma unroll +#pragma unroll for (int i = 0; i < v_per_thread; i++) { out[i] = static_cast(o[i]); } @@ -254,8 +254,8 @@ template sums += o_offset * blocks + block_idx; maxs += o_offset * blocks + block_idx; - // Read the query - #pragma unroll +// Read the query +#pragma unroll for (int i = 0; i < qk_per_thread; i++) { q[i] = static_cast(scale) * queries[i]; } @@ -280,7 +280,7 @@ template if (use_key) { // Compute the i-th score U score = 0; - #pragma unroll +#pragma unroll for (int i = 0; i < qk_per_thread; i++) { score += q[i] * keys[i]; } @@ -298,8 +298,8 @@ template max_score = new_max; sum_exp_score = sum_exp_score * factor + exp_score; - // Update the output accumulator - #pragma unroll +// Update the output accumulator +#pragma unroll for (int i = 0; i < v_per_thread; i++) { o[i] = o[i] * factor + exp_score * values[i]; } @@ -322,7 +322,7 @@ template maxs[0] = max_score; } - #pragma unroll +#pragma unroll for (int i = 0; i < v_per_thread; i++) { out[i] = o[i]; } @@ -571,8 +571,8 @@ template for (int b = 0; b < blocks / BN; ++b) { U factor = fast::exp(maxs[simd_gid] - max_score); - // Update the output accumulator - #pragma unroll +// Update the output accumulator +#pragma unroll for (int i = 0; i < elem_per_thread; i++) { o[i] += factor * static_cast(partials[i]); } @@ -581,8 +581,8 @@ template partials += BN * D; } - // Use shared memory to transpose and reduce the final block - #pragma unroll +// Use shared memory to transpose and reduce the final block +#pragma unroll for (int i = 0; i < elem_per_thread; i++) { outputs[simd_lid * BD + simd_gid] = o[i]; threadgroup_barrier(mem_flags::mem_threadgroup); @@ -593,7 +593,7 @@ template // And write the output if (simd_lid == 0) { - #pragma unroll +#pragma unroll for (int i = 0; i < elem_per_thread; i++) { out[i] = static_cast(o[i]); } diff --git a/mlx/backend/metal/kernels/softmax.h b/mlx/backend/metal/kernels/softmax.h index d995610649..cbbdbb41a2 100644 --- a/mlx/backend/metal/kernels/softmax.h +++ b/mlx/backend/metal/kernels/softmax.h @@ -27,12 +27,12 @@ template in += gid * size_t(axis_size) + lid * N_READS; if (lid * N_READS + N_READS <= axis_size) { - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { ld[i] = AccT(in[i]); } } else { - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { ld[i] = ((lid * N_READS + i) < axis_size) ? AccT(in[i]) : Limits::min; @@ -46,7 +46,7 @@ template // Get the max AccT maxval = Limits::finite_min; - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { maxval = (maxval < ld[i]) ? ld[i] : maxval; } @@ -66,7 +66,7 @@ template // Compute exp(x_i - maxval) and store the partial sums in local_normalizer AccT normalizer = 0; - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { AccT exp_x = softmax_exp(ld[i] - maxval); ld[i] = exp_x; @@ -89,12 +89,12 @@ template // Normalize and write to the output out += gid * size_t(axis_size) + lid * N_READS; if (lid * N_READS + N_READS <= axis_size) { - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { out[i] = T(ld[i] * normalizer); } } else { - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { if ((lid * N_READS + i) < axis_size) { out[i] = T(ld[i] * normalizer); @@ -129,24 +129,24 @@ template int offset = r * lsize * N_READS + lid * N_READS; AccT vals[N_READS]; if (offset + N_READS <= axis_size) { - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { vals[i] = AccT(in[offset + i]); } } else { - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { vals[i] = (offset + i < axis_size) ? AccT(in[offset + i]) : Limits::min; } } prevmax = maxval; - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { maxval = (maxval < vals[i]) ? vals[i] : maxval; } normalizer *= softmax_exp(prevmax - maxval); - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { normalizer += softmax_exp(vals[i] - maxval); } @@ -185,12 +185,12 @@ template r++) { int offset = r * lsize * N_READS + lid * N_READS; if (offset + N_READS <= axis_size) { - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { out[offset + i] = T(softmax_exp(in[offset + i] - maxval) * normalizer); } } else { - #pragma unroll +#pragma unroll for (int i = 0; i < N_READS; i++) { if (offset + i < axis_size) { out[offset + i] = diff --git a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h index 46c4f35f11..9176928669 100644 --- a/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h +++ b/mlx/backend/metal/kernels/steel/gemm/kernels/steel_gemm_fused.h @@ -211,41 +211,42 @@ template < // (complex64) never instantiate the branch. bool synthesized = false; if constexpr (kCausalBiasSynthEligible) { - if (addmm_params->fdc == 1 && - addmm_params->ldc == params->N + 1 && params->M <= params->N && - c_bstride_zero) { - // CAUSAL-CLOAD (concept receipt: solver i34-9, submission d0ccbe3c). - // A row stride of N + 1 on a bf16 addmm source operand cannot arise - // from any contiguous or broadcast operand of the declared output - // width; it is the deliberate signature of the composed-prefill - // causal bias view, and of nothing else. Synthesize that operand's - // two constants instead of loading them: widened bfloat16 lowest - // finite (0xFF7F) strictly above the causal diagonal placed at - // N - M, widened bfloat16 negative zero on and below it. The addend - // enters through the same TransformAdd, per accumulator element, with - // the same widening as the loaded operand it replaces, so every - // stored word is bit-identical. The padded backing store keeps every - // non-synthesizing branch load-correct at this stride. - const int diag = params->N - params->M; - const int row0 = c_row + mma_op.sm; - const int col0 = c_col + mma_op.sn; - const AccumType mask_add = - static_cast(as_type(0xFF7F0000u)); - const AccumType pass_add = static_cast(-0.0f); - STEEL_PRAGMA_UNROLL - for (short i = 0; i < mma_t::TM; i++) { + if (addmm_params->fdc == 1 && addmm_params->ldc == params->N + 1 && + params->M <= params->N && c_bstride_zero) { + // CAUSAL-CLOAD (concept receipt: solver i34-9, submission + // d0ccbe3c). A row stride of N + 1 on a bf16 addmm source operand + // cannot arise from any contiguous or broadcast operand of the + // declared output width; it is the deliberate signature of the + // composed-prefill causal bias view, and of nothing else. + // Synthesize that operand's two constants instead of loading them: + // widened bfloat16 lowest finite (0xFF7F) strictly above the causal + // diagonal placed at N - M, widened bfloat16 negative zero on and + // below it. The addend enters through the same TransformAdd, per + // accumulator element, with the same widening as the loaded operand + // it replaces, so every stored word is bit-identical. The padded + // backing store keeps every non-synthesizing branch load-correct at + // this stride. + const int diag = params->N - params->M; + const int row0 = c_row + mma_op.sm; + const int col0 = c_col + mma_op.sn; + const AccumType mask_add = + static_cast(as_type(0xFF7F0000u)); + const AccumType pass_add = static_cast(-0.0f); STEEL_PRAGMA_UNROLL - for (short j = 0; j < mma_t::TN; j++) { - thread auto& accum = mma_op.Ctile.frag_at(i, j); - const int row = row0 + i * mma_t::TM_stride; - const int col = col0 + j * mma_t::TN_stride; + for (short i = 0; i < mma_t::TM; i++) { STEEL_PRAGMA_UNROLL - for (short k = 0; k < decltype(mma_op.Ctile)::kElemsPerFrag; k++) { - accum[k] = epilogue_op_add.apply( - accum[k], (col + k) - row <= diag ? pass_add : mask_add); + for (short j = 0; j < mma_t::TN; j++) { + thread auto& accum = mma_op.Ctile.frag_at(i, j); + const int row = row0 + i * mma_t::TM_stride; + const int col = col0 + j * mma_t::TN_stride; + STEEL_PRAGMA_UNROLL + for (short k = 0; k < decltype(mma_op.Ctile)::kElemsPerFrag; + k++) { + accum[k] = epilogue_op_add.apply( + accum[k], (col + k) - row <= diag ? pass_add : mask_add); + } } } - } synthesized = true; } } diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index 47a1a4833a..0c95d41510 100644 --- a/mlx/backend/metal/scaled_dot_product_attention.cpp +++ b/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -52,8 +52,7 @@ inline bool env_flag_on(const char* name, bool default_on) { } inline bool d512_vector_sdpa_enabled() { - static bool enabled = - env_flag_on("DARKBLOOM_GEMMA4_D512_DECODE_2PASS", true); + static bool enabled = env_flag_on("DARKBLOOM_GEMMA4_D512_DECODE_2PASS", true); return enabled; }