diff --git a/fork.yaml b/fork.yaml index 0bf5f5bc56..5e3dafd4cb 100644 --- a/fork.yaml +++ b/fork.yaml @@ -119,6 +119,36 @@ 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" + - "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 17e74d1db1..2dd41bc90e 100644 --- a/mlx/backend/metal/kernels/quantized.h +++ b/mlx/backend/metal/kernels/quantized.h @@ -289,6 +289,178 @@ 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 +993,397 @@ 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 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. +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 += + float(xm[i]) + float(xm[i + 1]) + float(xm[i + 2]) + float(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++) { @@ -949,37 +1512,901 @@ METAL_FUNC void qmv_impl( 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 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); + 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))); +} - 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; +// Textual twin of `load_vector`'s `sum` on the same aligned +// 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]; + 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 += 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; +} - U s = sl[0]; - U b = bl[0]; - result[row] += qdot_safe( - wl, x_thread, s, b, sum, remaining); - } +#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; + } + + 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 +3177,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 +3199,260 @@ 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 +3490,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 +3512,184 @@ 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 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 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 +4135,321 @@ 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 +4477,232 @@ 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( + 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, diff --git a/mlx/backend/metal/kernels/quantized_nax.h b/mlx/backend/metal/kernels/quantized_nax.h index ed32eb59a7..4aa5fdf254 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,43 @@ 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 +1810,42 @@ 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_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; @@ -1648,9 +1883,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; @@ -1679,20 +1917,20 @@ 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)); } }); }); diff --git a/mlx/backend/metal/kernels/scaled_dot_product_attention.metal b/mlx/backend/metal/kernels/scaled_dot_product_attention.metal index cfe5aec6c3..a5f132507f 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_fp32partials_1_gqa_" #type "_" #qk_dim "_" #value_dim, \ sdpa_vector_2pass_1_gqa, \ @@ -36,7 +40,18 @@ using namespace metal; qk_dim, \ value_dim, \ 8, \ - hpt) + hpt, \ + split) + +// 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, \ + sdpa_vector_2pass_1, \ + type, \ + qk_dim, \ + value_dim) #define instantiate_sdpa_vector_heads(type) \ instantiate_sdpa_vector(type, 64, 64) \ @@ -45,13 +60,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_gqa(type, 64, 64, 8) \ - instantiate_sdpa_vector_gqa(type, 128, 128, 4) \ + instantiate_sdpa_vector_2pass(type, 512, 512) \ + 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) \ 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/kernels/sdpa_vector.h b/mlx/backend/metal/kernels/sdpa_vector.h index 3631e49daa..6c73ea5d8a 100644 --- a/mlx/backend/metal/kernels/sdpa_vector.h +++ b/mlx/backend/metal/kernels/sdpa_vector.h @@ -80,10 +80,12 @@ template out += o_offset * V + simd_gid * v_per_thread; - // Read the query and 0 the output accumulator +// 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; } @@ -106,13 +108,15 @@ template use_key = (fmask[0] >= Limits::finite_min); } if (use_key) { - // Read the 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]; } @@ -129,7 +133,8 @@ template max_score = new_max; sum_exp_score = sum_exp_score * factor + exp_score; - // Update the output accumulator +// Update the output accumulator +#pragma unroll for (int j = 0; j < v_per_thread; j++) { o[j] = o[j] * factor + exp_score * values[j]; } @@ -159,7 +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 +// 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]); } @@ -247,7 +254,8 @@ template sums += o_offset * blocks + block_idx; maxs += o_offset * blocks + block_idx; - // Read the query +// 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]; } @@ -289,7 +298,8 @@ template max_score = new_max; sum_exp_score = sum_exp_score * factor + exp_score; - // Update the output accumulator +// 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] = o[i]; } @@ -322,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)]], @@ -418,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]; + } } } @@ -510,7 +571,8 @@ template for (int b = 0; b < blocks / BN; ++b) { U factor = fast::exp(maxs[simd_gid] - max_score); - // Update the output accumulator +// Update the output accumulator +#pragma unroll for (int i = 0; i < elem_per_thread; i++) { o[i] += factor * static_cast(partials[i]); } @@ -519,7 +581,8 @@ template partials += BN * D; } - // Use shared memory to transpose and reduce the final block +// 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 +593,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..cbbdbb41a2 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] = 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..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 @@ -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,54 @@ 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)); diff --git a/mlx/backend/metal/scaled_dot_product_attention.cpp b/mlx/backend/metal/scaled_dot_product_attention.cpp index c2ee37f29e..0c95d41510 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,83 @@ 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. +// +// 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); + 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 +547,9 @@ void sdpa_vector_2pass( kname.reserve(64); kname += "sdpa_vector_2pass_fp32partials_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"; } @@ -684,7 +766,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 " @@ -873,7 +956,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 { diff --git a/python/tests/test_fast_sdpa.py b/python/tests/test_fast_sdpa.py index 3a9c04b075..56afcd42a7 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,56 @@ 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): + 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) + 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 + ) + 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 +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, 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") + 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 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