Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 30 additions & 0 deletions fork.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading
Loading