diff --git a/crates/metaltile-codegen/src/spirv/mod.rs b/crates/metaltile-codegen/src/spirv/mod.rs index a3a2ea3a..ff1142af 100644 --- a/crates/metaltile-codegen/src/spirv/mod.rs +++ b/crates/metaltile-codegen/src/spirv/mod.rs @@ -2098,7 +2098,7 @@ pub fn safe_glsl_ident(name: &str) -> String { "image3D", // GLSL reserves these for future use, even though they aren't // currently used; shaderc rejects them. The corpus hit them on - // conv* / depthwise_conv2d / aura_value_int4 / ffai_gemm kernels. + // conv* / depthwise_conv2d / aura_value_int4 / mt_gemm kernels. "input", "output", "texture", diff --git a/crates/metaltile-std/src/ffai/ffai_dequant_q4.rs b/crates/metaltile-std/src/ffai/ffai_dequant_q4.rs index 822a5323..b3dc9fc4 100644 --- a/crates/metaltile-std/src/ffai/ffai_dequant_q4.rs +++ b/crates/metaltile-std/src/ffai/ffai_dequant_q4.rs @@ -7,7 +7,7 @@ //! the inline-dequant `gemv_q4`, which is bandwidth-bound — perfect at //! batch 1. Prefill is COMPUTE-bound: many tokens reuse each weight, so //! the right move is to dequant the weight ONCE to f16 and feed the -//! tensor-core `ffai_gemm`. This kernel is that one-time expansion. +//! tensor-core `mt_gemm`. This kernel is that one-time expansion. //! //! ## Q4 block layout (matches `ffai_ops::quantize_q4`) //! diff --git a/crates/metaltile-std/src/ffai/gemm_q4_mpp.rs b/crates/metaltile-std/src/ffai/gemm_q4_mpp.rs index 3c1a87fe..c140cfd8 100644 --- a/crates/metaltile-std/src/ffai/gemm_q4_mpp.rs +++ b/crates/metaltile-std/src/ffai/gemm_q4_mpp.rs @@ -5,7 +5,7 @@ //! bench's Q4 weight layout (signed 4-bit, per-32-block scale `amax/7` stored //! f16) instead of Q8_0. Compute-bound prefill projections (q/k/v/o, mamba //! in/out_proj, shared experts, lm_head) run on this instead of the f32 -//! scalar `ffai_gemm` (which sat at ~0.1% of the tensor-core peak). +//! scalar `mt_gemm` (which sat at ~0.1% of the tensor-core peak). //! //! Same 64×64×32 coop_tile geometry as `ffai_gemm_q8_mpp` (4 simdgroups, //! 2×2 warp grid, 128 threads/tg). Only the weight-dequant block differs. diff --git a/crates/metaltile-std/src/ffai/gemm_q8.rs b/crates/metaltile-std/src/ffai/gemm_q8.rs index 9f77c9d5..9d473106 100644 --- a/crates/metaltile-std/src/ffai/gemm_q8.rs +++ b/crates/metaltile-std/src/ffai/gemm_q8.rs @@ -10,10 +10,10 @@ //! //! Weight is the resident Q8 split (`qs` int8 packed 4/u32 + per-32-block //! `d` scale), laid out as `[out_dim, in_dim]` (row-major over values). -//! Mirrors `ffai_gemm`'s geometry exactly so the dispatch wrapper is +//! Mirrors `mt_gemm`'s geometry exactly so the dispatch wrapper is //! identical apart from the two extra weight buffers. //! -//! ## DISPATCH INVARIANTS (same as ffai_gemm) +//! ## DISPATCH INVARIANTS (same as mt_gemm) //! - TPG = 1024 (32×32). Grid: (out_dim/32) × (n_rows/32) threadgroups. //! - `in_dim % 16 == 0` (K-tile) AND `in_dim % 32 == 0` (Q8 block). //! - Row/col edges handled in-kernel (clamp loads to 0, skip OOB stores). diff --git a/crates/metaltile-std/src/ffai/mod.rs b/crates/metaltile-std/src/ffai/mod.rs index 1b200326..7ea71518 100644 --- a/crates/metaltile-std/src/ffai/mod.rs +++ b/crates/metaltile-std/src/ffai/mod.rs @@ -59,11 +59,9 @@ pub mod gated_delta_prep_chunk; pub mod gated_delta_replay; pub mod gated_delta_wy; pub mod gelu_erf; -pub mod gemm; pub mod gemm_q4_mpp; pub mod gemm_q8; pub mod gemm_q8_mpp; -pub mod gemv_axpy_inplace; pub mod gemv_q8; pub mod gguf_dequant_iq2_xxs; pub mod gguf_dequant_iq2_xxs_raw; @@ -108,9 +106,7 @@ pub mod moe_mpp_int8; pub mod moe_mpp_shared; pub mod moe_router_sigmoid_bias; pub mod moe_router_sqrtsoftplus; -pub mod patch_embed; pub mod patch_embed_block_scaled; -pub mod patch_embed_mma; pub mod patch_embed_mma_block_scaled; pub mod patch_unfold_qwen; pub mod pos_emb_2d_add; diff --git a/crates/metaltile-std/src/ffai/ssm.rs b/crates/metaltile-std/src/ffai/ssm.rs index cba55a9b..c1675172 100644 --- a/crates/metaltile-std/src/ffai/ssm.rs +++ b/crates/metaltile-std/src/ffai/ssm.rs @@ -361,8 +361,8 @@ pub mod kernel_tests { }; let kernels: Vec<(&str, metaltile::core::Kernel)> = vec![ - ("ffai_gemm_batched", { - let mut k = super::super::gemm::ffai_gemm_batched::kernel_ir_for(DType::F32); + ("mt_gemm_batched", { + let mut k = crate::kernels::gemm::dense::mt_gemm_batched::kernel_ir_for(DType::F32); k.mode = KernelMode::Reduction; k }), diff --git a/crates/metaltile-std/src/ffai/gemm.rs b/crates/metaltile-std/src/kernels/gemm/dense.rs similarity index 93% rename from crates/metaltile-std/src/ffai/gemm.rs rename to crates/metaltile-std/src/kernels/gemm/dense.rs index 1c08461c..73fffad5 100644 --- a/crates/metaltile-std/src/ffai/gemm.rs +++ b/crates/metaltile-std/src/kernels/gemm/dense.rs @@ -34,7 +34,7 @@ use metaltile::kernel; #[kernel] -pub fn ffai_gemm( +pub fn mt_gemm( weight: Tensor, input: Tensor, out: Tensor, @@ -91,7 +91,7 @@ pub fn ffai_gemm( // // Computes, for each batch `z ∈ [0, batch)`: // out[z][r, o] = Σ_k weight[z][o, k] · input[z][r, k] -// i.e. the SAME `out = input · weightᵀ` contraction as `ffai_gemm`, replicated +// i.e. the SAME `out = input · weightᵀ` contraction as `mt_gemm`, replicated // over a batch axis. Each operand is a contiguous stack of per-batch matrices // with element strides `w_stride` / `x_stride` / `o_stride` (counted in T // elements, NOT bytes) — the wrapper passes `out_dim*in_dim`, `n_rows*in_dim`, @@ -100,15 +100,15 @@ pub fn ffai_gemm( // This is what the Mamba2 SSD chunked-matmul prefill scan runs its 4 batched // GEMMs on (batch = n_chunks·n_heads) on Apple/HIP/Vulkan hardware MMA, where // cuBLAS strided-batched is unavailable. Same 32×32 / 16-K tiling as -// `ffai_gemm`; the batch index rides the 3rd grid axis (`tgid_z`). +// `mt_gemm`; the batch index rides the 3rd grid axis (`tgid_z`). // -// ## DISPATCH INVARIANTS (identical to `ffai_gemm` plus the batch axis) +// ## DISPATCH INVARIANTS (identical to `mt_gemm` plus the batch axis) // - TPG = 1024 (BM·BN). `in_dim % 16 == 0`. // - Grid: `((out_dim+31)/32, (n_rows+31)/32, batch)` Reduction-mode TGs. // - All operands row-major within each batch slice; out-of-range row/col edges // handled in-kernel (clamp-load to 0, skip-store). #[kernel] -pub fn ffai_gemm_batched( +pub fn mt_gemm_batched( weight: Tensor, input: Tensor, out: Tensor, @@ -163,7 +163,7 @@ pub fn ffai_gemm_batched( } } -/// New-syntax correctness tests for `ffai_gemm` — the multi-row 32×32-tiled +/// New-syntax correctness tests for `mt_gemm` — the multi-row 32×32-tiled /// GEMM `out[r, :] = weight · input[r, :]`. Reduction-mode (threadgroup-memory /// tiles + barriers). /// @@ -179,7 +179,7 @@ pub fn ffai_gemm_batched( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::ffai_gemm; + use super::mt_gemm; use crate::utils::{pack_f32, unpack_f32}; /// Triple-loop reference: out[r, o] = Σ_k weight[o, k] · input[r, k]. @@ -211,7 +211,7 @@ pub mod kernel_tests { let w = unpack_f32(&pack_f32(&weight_f, dt), dt); let x = unpack_f32(&pack_f32(&input_f, dt), dt); let expected = gemm_oracle(&w, &x, n_rows, in_dim, out_dim); - TestSetup::new(ffai_gemm::kernel_ir_for(dt)) + TestSetup::new(mt_gemm::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .input(TestBuffer::from_vec("weight", pack_f32(&weight_f, dt), dt)) .input(TestBuffer::from_vec("input", pack_f32(&input_f, dt), dt)) @@ -231,8 +231,8 @@ pub mod kernel_tests { #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_gemm_edge(dt: DType) -> TestSetup { gemm_setup(20, 48, 100, dt) } - // ── ffai_gemm_batched ──────────────────────────────────────────────── - use super::ffai_gemm_batched; + // ── mt_gemm_batched ──────────────────────────────────────────────── + use super::mt_gemm_batched; fn gemm_batched_setup( batch: usize, @@ -257,7 +257,7 @@ pub mod kernel_tests { let eb = gemm_oracle(wb, xb, n_rows, in_dim, out_dim); expected[b * n_rows * out_dim..(b + 1) * n_rows * out_dim].copy_from_slice(&eb); } - TestSetup::new(ffai_gemm_batched::kernel_ir_for(dt)) + TestSetup::new(mt_gemm_batched::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .input(TestBuffer::from_vec("weight", pack_f32(&weight_f, dt), dt)) .input(TestBuffer::from_vec("input", pack_f32(&input_f, dt), dt)) @@ -283,19 +283,19 @@ pub mod kernel_tests { fn test_gemm_batched_edge(dt: DType) -> TestSetup { gemm_batched_setup(2, 20, 48, 100, dt) } } -/// New-syntax benchmark for `ffai_gemm`. Nemotron-class block-diffusion shape: +/// New-syntax benchmark for `mt_gemm`. Nemotron-class block-diffusion shape: /// a 32-row block projected through a `[out_dim, in_dim]` weight (hidden 4096). pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_gemm; + use super::mt_gemm; #[bench(dtypes = [f32, f16, bf16])] fn bench_gemm(dt: DType) -> BenchSetup { let (n_rows, in_dim, out_dim) = (32usize, 4096usize, 4096usize); let sz = dt.size_bytes(); let bytes = out_dim * in_dim * sz + n_rows * in_dim * sz + n_rows * out_dim * sz; - BenchSetup::new(ffai_gemm::kernel_ir_for(dt)) + BenchSetup::new(mt_gemm::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("weight", out_dim * in_dim, dt)) .buffer(BenchBuffer::random("input", n_rows * in_dim, dt)) diff --git a/crates/metaltile-std/src/mlx/gemv.rs b/crates/metaltile-std/src/kernels/gemm/gemv.rs similarity index 100% rename from crates/metaltile-std/src/mlx/gemv.rs rename to crates/metaltile-std/src/kernels/gemm/gemv.rs diff --git a/crates/metaltile-std/src/ffai/gemv_axpy_inplace.rs b/crates/metaltile-std/src/kernels/gemm/gemv_axpy_inplace.rs similarity index 94% rename from crates/metaltile-std/src/ffai/gemv_axpy_inplace.rs rename to crates/metaltile-std/src/kernels/gemm/gemv_axpy_inplace.rs index f1691cba..a71372a1 100644 --- a/crates/metaltile-std/src/ffai/gemv_axpy_inplace.rs +++ b/crates/metaltile-std/src/kernels/gemm/gemv_axpy_inplace.rs @@ -20,7 +20,7 @@ use metaltile::kernel; #[kernel] -pub fn ffai_gemv_axpy_inplace( +pub fn mt_gemv_axpy_inplace( mat: Tensor, vec: Tensor, mut accum: Tensor, @@ -39,7 +39,7 @@ pub fn ffai_gemv_axpy_inplace( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::ffai_gemv_axpy_inplace; + use super::mt_gemv_axpy_inplace; use crate::utils::{pack_f32, unpack_f32}; fn setup(m: usize, k: usize, weight: f32, dt: DType) -> TestSetup { @@ -55,7 +55,7 @@ pub mod kernel_tests { accum_dt[r] + weight * dot }) .collect(); - TestSetup::new(ffai_gemv_axpy_inplace::kernel_ir_for(dt)) + TestSetup::new(mt_gemv_axpy_inplace::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .input(TestBuffer::from_vec("mat", pack_f32(&mat, dt), dt)) .input(TestBuffer::from_vec("vec", pack_f32(&vec, dt), dt)) @@ -84,12 +84,12 @@ pub mod kernel_tests { pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_gemv_axpy_inplace; + use super::mt_gemv_axpy_inplace; #[bench(dtypes = [f32, f16, bf16])] fn bench_gemv_axpy(dt: DType) -> BenchSetup { let (m, k) = (4096usize, 2048usize); - BenchSetup::new(ffai_gemv_axpy_inplace::kernel_ir_for(dt)) + BenchSetup::new(mt_gemv_axpy_inplace::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("mat", m * k, dt)) .buffer(BenchBuffer::random("vec", k, dt)) diff --git a/crates/metaltile-std/src/mlx/gemv_masked.rs b/crates/metaltile-std/src/kernels/gemm/gemv_masked.rs similarity index 100% rename from crates/metaltile-std/src/mlx/gemv_masked.rs rename to crates/metaltile-std/src/kernels/gemm/gemv_masked.rs diff --git a/crates/metaltile-std/src/kernels/gemm/mod.rs b/crates/metaltile-std/src/kernels/gemm/mod.rs new file mode 100644 index 00000000..ddb53108 --- /dev/null +++ b/crates/metaltile-std/src/kernels/gemm/mod.rs @@ -0,0 +1,20 @@ +//! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric +//! SPDX-License-Identifier: Apache-2.0 +//! Matrix-multiply kernels — the gemm family (see +//! `docs/specs/KERNEL_CONSOLIDATION_PLAN.md`): dense GEMM / GEMV (+ masked, +//! +axpy), the batched QKV / 4-way projection forms, patch-embed (im2col + +//! matmul), and the MLX `steel/` tiled-GEMM templates. +//! +//! This is the **dense** half of the family, migrated from the legacy `mlx/` + +//! `ffai/` split. The quantized matmuls (`*_q8`/`q4`/`block_scaled_qmm`/ +//! `quantized_*`/`fp_quantized_*` and the batched `*_qgemv`/`*_qmm` forms) are +//! the quantized form of these ops and land here in a follow-up pass; the +//! format-axis fold (plan §7) is deferred. + +pub mod dense; +pub mod gemv; +pub mod gemv_axpy_inplace; +pub mod gemv_masked; +pub mod patch_embed; +pub mod patch_embed_mma; +pub mod steel; diff --git a/crates/metaltile-std/src/ffai/patch_embed.rs b/crates/metaltile-std/src/kernels/gemm/patch_embed.rs similarity index 96% rename from crates/metaltile-std/src/ffai/patch_embed.rs rename to crates/metaltile-std/src/kernels/gemm/patch_embed.rs index 78e16324..3440b946 100644 --- a/crates/metaltile-std/src/ffai/patch_embed.rs +++ b/crates/metaltile-std/src/kernels/gemm/patch_embed.rs @@ -14,7 +14,7 @@ //! the image and dots them with one weight row, no intermediate buffer. //! //! It differs from `conv2d` in layout, not arithmetic — `conv2d` keeps -//! the NCHW image convention and writes NCHW output; `patch_embed` takes +//! the NCHW image convention and writes NCHW output; `mt_patch_embed` takes //! the same NCHW image but treats the weight as a flat linear matrix //! `[hidden, patch_dim]` and writes transformer-token output //! `[num_patches, hidden]`, which is what a ViT block consumes directly. @@ -46,7 +46,7 @@ use metaltile::kernel; #[kernel] -pub fn patch_embed( +pub fn mt_patch_embed( image: Tensor, weight: Tensor, bias: Tensor, @@ -93,7 +93,7 @@ pub fn patch_embed( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::patch_embed; + use super::mt_patch_embed; use crate::utils::{pack_f32, unpack_f32}; fn ramp(n: usize, period: usize, amp: f32) -> Vec { @@ -164,7 +164,7 @@ pub mod kernel_tests { let bias = unpack_f32(&pack_f32(&bias_f, dt), dt); let expected = naive_patch_embed(&image, &weight, &bias, in_ch, in_h, in_w, patch_h, patch_w, hidden); - TestSetup::new(patch_embed::kernel_ir_for(dt)) + TestSetup::new(mt_patch_embed::kernel_ir_for(dt)) .mode(KernelMode::Grid3D) .input(TestBuffer::from_vec("image", pack_f32(&image_f, dt), dt)) .input(TestBuffer::from_vec("weight", pack_f32(&weight_f, dt), dt)) @@ -199,12 +199,12 @@ pub mod kernel_tests { } } -/// New-syntax bench for `patch_embed` (ViT-L SigLIP stem shape). +/// New-syntax bench for `mt_patch_embed` (ViT-L SigLIP stem shape). /// Grid3D, `grid_1d(num_patches * hidden, 256)`; bytes_moved = output stream. pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::patch_embed; + use super::mt_patch_embed; #[bench(dtypes = [f32, f16, bf16])] fn bench_patch_embed(dt: DType) -> BenchSetup { @@ -215,7 +215,7 @@ pub mod kernel_benches { let num_patches = (in_h / patch_h) * (in_w / patch_w); let patch_dim = in_ch * patch_h * patch_w; let n_out = num_patches * hidden; - BenchSetup::new(patch_embed::kernel_ir_for(dt)) + BenchSetup::new(mt_patch_embed::kernel_ir_for(dt)) .mode(KernelMode::Grid3D) .buffer(BenchBuffer::random("image", in_ch * in_h * in_w, dt)) .buffer(BenchBuffer::random("weight", hidden * patch_dim, dt)) diff --git a/crates/metaltile-std/src/ffai/patch_embed_mma.rs b/crates/metaltile-std/src/kernels/gemm/patch_embed_mma.rs similarity index 98% rename from crates/metaltile-std/src/ffai/patch_embed_mma.rs rename to crates/metaltile-std/src/kernels/gemm/patch_embed_mma.rs index 8af992fa..be44b535 100644 --- a/crates/metaltile-std/src/ffai/patch_embed_mma.rs +++ b/crates/metaltile-std/src/kernels/gemm/patch_embed_mma.rs @@ -72,7 +72,7 @@ use metaltile::kernel; /// Correctness pinned by the in-source `#[test_kernel]`s. #[kernel] #[allow(clippy::too_many_arguments)] -pub fn patch_embed_mma( +pub fn mt_patch_embed_mma( image: Tensor, weight: Tensor, bias: Tensor, @@ -292,7 +292,7 @@ pub fn patch_embed_mma( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::patch_embed_mma; + use super::mt_patch_embed_mma; use crate::utils::{pack_f32, unpack_f32}; fn ramp(n: usize, period: usize, amp: f32) -> Vec { @@ -367,7 +367,7 @@ pub mod kernel_tests { let expected = naive_patch_embed_mma( &image, &weight, &bias, in_ch, in_h, in_w, patch_h, patch_w, hidden, ); - TestSetup::new(patch_embed_mma::kernel_ir_for(dt)) + TestSetup::new(mt_patch_embed_mma::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .input(TestBuffer::from_vec("image", pack_f32(&image_f, dt), dt)) .input(TestBuffer::from_vec("weight", pack_f32(&weight_f, dt), dt)) @@ -392,12 +392,12 @@ pub mod kernel_tests { fn test_patch_embed_mma_4x4(dt: DType) -> TestSetup { mma_setup(8, 32, 32, 4, 4, 32, dt) } } -/// New-syntax bench for `patch_embed_mma` (ViT-L-ish 8×8 patch, hidden 1024). +/// New-syntax bench for `mt_patch_embed_mma` (ViT-L-ish 8×8 patch, hidden 1024). /// Reduction mode, `grid_3d(hidden/32, num_patches/32, 1, [128,1,1])`. pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::patch_embed_mma; + use super::mt_patch_embed_mma; #[bench(dtypes = [f32, f16, bf16])] fn bench_patch_embed_mma(dt: DType) -> BenchSetup { @@ -408,7 +408,7 @@ pub mod kernel_benches { let num_patches = (in_h / patch_h) * (in_w / patch_w); let patch_dim = in_ch * patch_h * patch_w; let n_out = num_patches * hidden; - BenchSetup::new(patch_embed_mma::kernel_ir_for(dt)) + BenchSetup::new(mt_patch_embed_mma::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("image", in_ch * in_h * in_w, dt)) .buffer(BenchBuffer::random("weight", hidden * patch_dim, dt)) diff --git a/crates/metaltile-std/src/mlx/steel/gemm/mod.rs b/crates/metaltile-std/src/kernels/gemm/steel/mod.rs similarity index 100% rename from crates/metaltile-std/src/mlx/steel/gemm/mod.rs rename to crates/metaltile-std/src/kernels/gemm/steel/mod.rs diff --git a/crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_fused.rs b/crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_fused.rs similarity index 100% rename from crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_fused.rs rename to crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_fused.rs diff --git a/crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_fused_nax.rs b/crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_fused_nax.rs similarity index 100% rename from crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_fused_nax.rs rename to crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_fused_nax.rs diff --git a/crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_gather.rs b/crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_gather.rs similarity index 100% rename from crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_gather.rs rename to crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_gather.rs diff --git a/crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_gather_nax.rs b/crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_gather_nax.rs similarity index 100% rename from crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_gather_nax.rs rename to crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_gather_nax.rs diff --git a/crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_masked.rs b/crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_masked.rs similarity index 100% rename from crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_masked.rs rename to crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_masked.rs diff --git a/crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_segmented.rs b/crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_segmented.rs similarity index 100% rename from crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_segmented.rs rename to crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_segmented.rs diff --git a/crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_splitk.rs b/crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_splitk.rs similarity index 100% rename from crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_splitk.rs rename to crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_splitk.rs diff --git a/crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_splitk_nax.rs b/crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_splitk_nax.rs similarity index 100% rename from crates/metaltile-std/src/mlx/steel/gemm/steel_gemm_splitk_nax.rs rename to crates/metaltile-std/src/kernels/gemm/steel/steel_gemm_splitk_nax.rs diff --git a/crates/metaltile-std/src/kernels/mod.rs b/crates/metaltile-std/src/kernels/mod.rs index 89cac5c7..18443bc0 100644 --- a/crates/metaltile-std/src/kernels/mod.rs +++ b/crates/metaltile-std/src/kernels/mod.rs @@ -8,6 +8,7 @@ //! consumer is regenerated from the new inventory after each family lands. pub mod convolution; +pub mod gemm; pub mod norm; pub mod ops; pub mod rope; diff --git a/crates/metaltile-std/src/mlx/mod.rs b/crates/metaltile-std/src/mlx/mod.rs index eaffe58a..181fdda3 100644 --- a/crates/metaltile-std/src/mlx/mod.rs +++ b/crates/metaltile-std/src/mlx/mod.rs @@ -25,8 +25,6 @@ pub mod fft; pub mod fp_quantized; pub mod fp_quantized_mma; pub mod fp_quantized_nax; -pub mod gemv; -pub mod gemv_masked; pub mod quantized; pub mod quantized_mma_dynamic_m; pub mod quantized_mpp; diff --git a/crates/metaltile-std/src/mlx/steel/mod.rs b/crates/metaltile-std/src/mlx/steel/mod.rs index 1682daf7..efe6f405 100644 --- a/crates/metaltile-std/src/mlx/steel/mod.rs +++ b/crates/metaltile-std/src/mlx/steel/mod.rs @@ -2,4 +2,3 @@ //! SPDX-License-Identifier: Apache-2.0 pub mod attn; pub use crate::kernels::convolution::steel_conv as conv; -pub mod gemm; diff --git a/crates/metaltile-std/tests/gemm_q8_correctness.rs b/crates/metaltile-std/tests/gemm_q8_correctness.rs index 078f1afd..7fd16715 100644 --- a/crates/metaltile-std/tests/gemm_q8_correctness.rs +++ b/crates/metaltile-std/tests/gemm_q8_correctness.rs @@ -32,7 +32,7 @@ fn run_case(dt: Dt, in_dim: usize, out_dim: usize, n_rows: usize, tol: f32) { (0..n_rows * in_dim).map(|_| ((xorshift(&mut st) % 2000) as f32 / 1000.0) - 1.0).collect(); // Round inputs to the kernel's dtype so the f32 oracle sees the same // values the kernel loads (f16 input rounding, amplified by cancellation - // over in_dim terms, otherwise dominates the diff). Matches ffai_gemm's test. + // over in_dim terms, otherwise dominates the diff). Matches mt_gemm's test. if matches!(dt, Dt::F16) { input = unpack_bytes(&pack_bytes(&input, Dt::F16), Dt::F16); } diff --git a/crates/metaltile-std/tests/steel_msl_snapshots.rs b/crates/metaltile-std/tests/steel_msl_snapshots.rs index 9493ba83..b618fc06 100644 --- a/crates/metaltile-std/tests/steel_msl_snapshots.rs +++ b/crates/metaltile-std/tests/steel_msl_snapshots.rs @@ -33,12 +33,11 @@ use metaltile::{ core::{dtype::DType, ir::KernelMode}, }; use metaltile_std::{ - kernels::ops::hadamard_m, - mlx::{ - fp_quantized_nax, - quantized_nax, - steel::gemm::{steel_gemm_fused_nax, steel_gemm_gather_nax, steel_gemm_splitk_nax}, + kernels::{ + gemm::steel::{steel_gemm_fused_nax, steel_gemm_gather_nax, steel_gemm_splitk_nax}, + ops::hadamard_m, }, + mlx::{fp_quantized_nax, quantized_nax}, }; /// Lower one of the NAX-family kernel IRs to MSL with its declared diff --git a/docs/specs/KERNEL_AUDIT.md b/docs/specs/KERNEL_AUDIT.md index def4be87..e30c8145 100644 --- a/docs/specs/KERNEL_AUDIT.md +++ b/docs/specs/KERNEL_AUDIT.md @@ -43,9 +43,9 @@ The 7 NAX kernels: | `mt_qmm_nax` | `mlx/quantized_nax.rs` | int4 quantized matmul prefill | | `mt_qmm_nax_int8` | `mlx/quantized_nax_int8.rs` | int8 quantized matmul prefill | | `mt_fp_qmm_nax` | `mlx/fp_quantized_nax.rs` | fp4 (E2M1) quantized matmul prefill | -| `mt_steel_gemm_fused_nax` | `mlx/steel/gemm/steel_gemm_fused_nax.rs` | plain fused GEMM | -| `mt_steel_gemm_gather_nax` | `mlx/steel/gemm/steel_gemm_gather_nax.rs` | MoE gather GEMM | -| `mt_steel_gemm_splitk_nax` + `_accum_nax` | `mlx/steel/gemm/steel_gemm_splitk_nax.rs` | split-K GEMM (pass1 + pass2) | +| `mt_steel_gemm_fused_nax` | `kernels/gemm/steel/steel_gemm_fused_nax.rs` | plain fused GEMM | +| `mt_steel_gemm_gather_nax` | `kernels/gemm/steel/steel_gemm_gather_nax.rs` | MoE gather GEMM | +| `mt_steel_gemm_splitk_nax` + `_accum_nax` | `kernels/gemm/steel/steel_gemm_splitk_nax.rs` | split-K GEMM (pass1 + pass2) | | `mt_sdpa_prefill_nax` | `mlx/steel/attn/steel_attention_nax.rs` | FlashAttention-2 prefill | The `quantized_mpp` family (`mt_qmm_mma_mpp`, `mt_qmm_mma_mpp_int8`, the four MoE `*_mpp` variants) uses the same MPP cooperative-tensor primitive and is similarly runtime-gated via `skip_unless_apple10`. The distinction between `*_mpp` and `*_nax`: `quantized_mpp` and its MoE siblings have working MXU-fallback paths on M1–M3 via Apple's `matmul2d` itself (slower than NAX hardware but functionally correct), whereas the `*_nax` kernels were authored specifically to exercise the M4+ tensor-core descriptor and have no fallback. @@ -88,20 +88,20 @@ Local verification of NAX kernels is the developer's responsibility on M4+ hardw | steel_attention (Flash, prefill) | ✓ | ✓ | ✓ | `mlx/steel/attn/steel_attention.rs` → `mt_sdpa_prefill`. Scalar-flash prefill (BQ=4, online softmax, causal), generic `T`, head_dim=128. | | steel_attention_mma (Flash prefill, simdgroup-MMA) | ✓ | ✓ | ✓ | `mlx/steel/attn/steel_attention_mma.rs` → `mt_sdpa_prefill_mma`. Real simdgroup-matrix MMA path; head_dim=128. A pre-M3 bf16-tuned sibling (`steel_attention_mma_bf16.rs`) is selected by `sdpa_prefill_mma_for()`. | | steel_attention_nax | ✓ | ✓ | ✓ | `mlx/steel/attn/steel_attention_nax.rs` → `mt_sdpa_prefill_nax` (d=32 base) + `mt_sdpa_prefill_nax_d{64,128,256}`. Flash-attention prefill via Apple `mpp::tensor_ops::matmul2d`. The wide variants loop the QK contraction over `head_dim/32` consecutive 32-wide D-chunks inside the outer K-block loop (first chunk uses `overwrite` descriptor, subsequent chunks `accumulate`); PV stores each chunk to a scratch `Opv` tile then accumulates into the full-width O buffer. Causal masking + GQA. Runtime-gated to Apple10+. | -| steel_gemm_fused | ✓ | ✓ | ✓ | `mlx/steel/gemm/steel_gemm_fused.rs` → `mt_steel_gemm_{32x32x16_1x2,32x64x16_1x2,32x32x16_2x2,64x64x16_2x2}` (4 block shapes via `instantiate_gemm_shapes_helper`). | -| steel_gemm_fused_nax | ✓ | ✓ | ✓ | `mlx/steel/gemm/steel_gemm_fused_nax.rs` → `mt_steel_gemm_fused_nax`. Plain fused GEMM `C = A·B` via NAX cooperative-tensor `matmul2d`. Runtime-gated to Apple10+. | -| steel_gemm_gather | ✓ | ✓ | ✓ | `mlx/steel/gemm/steel_gemm_gather.rs` → `mt_steel_gemm_gather_{64x64x16_2x2,32x32x16_2x2}`. Row-major `C = A_gathered·B_gathered` (MLX `gather_mm`, the dense matmul of a MoE FFN). | -| steel_gemm_gather_nax | ✓ | ✓ | ✓ | `mlx/steel/gemm/steel_gemm_gather_nax.rs` → `mt_steel_gemm_gather_nax`. Gather GEMM via NAX `matmul2d`. Runtime-gated to Apple10+. | -| steel_gemm_masked | ✓ | ✓ | ✓ | `mlx/steel/gemm/steel_gemm_masked.rs` → `mt_steel_gemm_masked_{64x64x16_2x2,32x32x16_2x2}`. Block-masked `C = A·B` (output-block mask zeros whole `BM×BN` blocks; operand-block mask scales each K-block contribution). | -| steel_gemm_segmented | ✓ | ✓ | ✓ | `mlx/steel/gemm/steel_gemm_segmented.rs` → `mt_steel_gemm_segmented_{64x64x16_2x2,32x32x16_2x2}`. Ragged-K batched matmul (MLX `segmented_mm`); each segment sums over its own `[k_start, k_end)` range. | -| steel_gemm_splitk + accum | ✓ | ✓ | ✓ | `mlx/steel/gemm/steel_gemm_splitk.rs` → pass 1 `mt_steel_gemm_splitk_{64x64x16_2x2,32x32x16_2x2}` + pass 2 `mt_steel_gemm_splitk_accum` / `mt_steel_gemm_splitk_accum_axpby`. Partials stay fp32 for cross-split precision on f16/bf16 inputs. | -| steel_gemm_splitk_nax | ✓ | ✓ | ✓ | `mlx/steel/gemm/steel_gemm_splitk_nax.rs` → pass 1 `mt_steel_gemm_splitk_nax` + pass 2 `mt_steel_gemm_splitk_accum_nax`. Split-K via NAX `matmul2d`; partials fp32. Runtime-gated to Apple10+. | +| steel_gemm_fused | ✓ | ✓ | ✓ | `kernels/gemm/steel/steel_gemm_fused.rs` → `mt_steel_gemm_{32x32x16_1x2,32x64x16_1x2,32x32x16_2x2,64x64x16_2x2}` (4 block shapes via `instantiate_gemm_shapes_helper`). | +| steel_gemm_fused_nax | ✓ | ✓ | ✓ | `kernels/gemm/steel/steel_gemm_fused_nax.rs` → `mt_steel_gemm_fused_nax`. Plain fused GEMM `C = A·B` via NAX cooperative-tensor `matmul2d`. Runtime-gated to Apple10+. | +| steel_gemm_gather | ✓ | ✓ | ✓ | `kernels/gemm/steel/steel_gemm_gather.rs` → `mt_steel_gemm_gather_{64x64x16_2x2,32x32x16_2x2}`. Row-major `C = A_gathered·B_gathered` (MLX `gather_mm`, the dense matmul of a MoE FFN). | +| steel_gemm_gather_nax | ✓ | ✓ | ✓ | `kernels/gemm/steel/steel_gemm_gather_nax.rs` → `mt_steel_gemm_gather_nax`. Gather GEMM via NAX `matmul2d`. Runtime-gated to Apple10+. | +| steel_gemm_masked | ✓ | ✓ | ✓ | `kernels/gemm/steel/steel_gemm_masked.rs` → `mt_steel_gemm_masked_{64x64x16_2x2,32x32x16_2x2}`. Block-masked `C = A·B` (output-block mask zeros whole `BM×BN` blocks; operand-block mask scales each K-block contribution). | +| steel_gemm_segmented | ✓ | ✓ | ✓ | `kernels/gemm/steel/steel_gemm_segmented.rs` → `mt_steel_gemm_segmented_{64x64x16_2x2,32x32x16_2x2}`. Ragged-K batched matmul (MLX `segmented_mm`); each segment sums over its own `[k_start, k_end)` range. | +| steel_gemm_splitk + accum | ✓ | ✓ | ✓ | `kernels/gemm/steel/steel_gemm_splitk.rs` → pass 1 `mt_steel_gemm_splitk_{64x64x16_2x2,32x32x16_2x2}` + pass 2 `mt_steel_gemm_splitk_accum` / `mt_steel_gemm_splitk_accum_axpby`. Partials stay fp32 for cross-split precision on f16/bf16 inputs. | +| steel_gemm_splitk_nax | ✓ | ✓ | ✓ | `kernels/gemm/steel/steel_gemm_splitk_nax.rs` → pass 1 `mt_steel_gemm_splitk_nax` + pass 2 `mt_steel_gemm_splitk_accum_nax`. Split-K via NAX `matmul2d`; partials fp32. Runtime-gated to Apple10+. | | steel_conv 2D (implicit-GEMM) | ✓ | ✓ | ✓ | `ffai/conv2d.rs` → `conv2d_patch14` / `conv2d_patch16` / `conv2d_generic`. Direct conv (implicit im2col, one thread per output). **MMA-tiled perf path** (PR #157): `ffai/conv2d_mma.rs` → `conv2d_mma` — implicit-im2col + 4-SG 2×2 simdgroup-matrix MMA, 32×32 output tile (stride=1/dilation=1/pad=0, out_ch and n_pixels divisible by 32). | | steel_conv 3D | ✓ | ✓ | ✓ | `ffai/conv3d.rs` → `conv3d_generic` + `conv3d_grouped` (depthwise + dilation). 5D NCDHW / OIDHW. **MMA-tiled perf path** (PR #157): `ffai/conv3d_mma.rs` → `conv3d_mma` — same MMA scaffold as 2D, decomposed over `(kd, kh, kw, ic)`. | | steel_conv_general (strides/dilation/groups) | ✓ | ✓ | ✓ | `ffai/conv2d.rs` → `conv2d_grouped`. Fully general 2D conv: strides, dilation (atrous), padding, grouped channels. | | conv (winograd + naive_unfold + depthwise) | ✓ | ✓ | ✓ | `ffai/conv2d.rs` / `ffai/conv3d.rs` cover `naive_unfold` + depthwise (via `_generic` / `_grouped` for both 2D and 3D). Winograd fast-conv: `ffai/winograd_conv.rs` → `winograd_conv2d_3x3` (F(2×2, 3×3) minimal-filtering, one thread per 2×2 output tile, requires even output dims) + `winograd_filter_transform_3x3` + `winograd_conv2d_3x3_split` (pre-transformed filters, removes O(tiles) redundant transform). | -| gemv | ✓ | ✓ | ✓ | `mlx/gemv.rs` → `mt_gemv`. | -| gemv_masked | ✓ | ✓ | ✓ | `mlx/gemv_masked.rs` → `mt_gemv_masked`. | +| gemv | ✓ | ✓ | ✓ | `kernels/gemm/gemv.rs` → `mt_gemv`. | +| gemv_masked | ✓ | ✓ | ✓ | `kernels/gemm/gemv_masked.rs` → `mt_gemv_masked`. | | quantized (affine_quantize / affine_dequantize) | ✓ | ✓ | ✓ | `mlx/quantized.rs` — quantize + dequantize for all widths: int2/int4/int8 (pack-aligned) + int3/int5/int6 (byte-stream). 12 kernels (`mt_affine_{quantize,dequantize}_int{2,3,4,5,6,8}`). int3/5/6 quantize uses bit-stream OR (lane 0 ORs codes into u32 words) to handle straddling — no atomics. | | quantized (affine_qmv / qvm / qmm — matvec / matmul) | ✓ | ✓ | ✓ | `mlx/quantized.rs` — **int4 perf**: `mt_qmv` (8-row-per-TG decode, mirrors MLX `qmv_fast`) + `mt_qmm` / `_bm2` / `_bm4` (M-batched prefill) + `mt_qmm_mma` / `_m16` (simdgroup-matrix MMA prefill) + `mt_qmm_mma_mpp` (MPP) + `mt_qmm_nax` (NAX). **int8 perf** (PR #154): `mt_qmv_int8_fast`, `mt_qmm_int8_fast` / `_bm2` / `_bm4`, `mt_qmm_mma_int8` / `_m16_int8`, `mt_qmm_mma_mpp_int8`, `mt_qmm_nax_int8` — pack-aligned (4 bytes/u32, byte-shift extract), closes the ~6–8× int8-vs-int4 perf gap. **Odd-bitwidth MMA** (PR #157): `mt_qmm_mma_b{3,5,6}` — straddle-aware two-word bit-stream dequant in the 4-SG MMA body. **All bit-widths × all dtypes**: `mt_{qmv,qvm,qmm}_b{3,4,5,6,8}` (correctness-first scalar family). **qvm perf**: `mt_qvm_int4_fast` (PR #154) — 8-col-per-TG, MLX `qvm_fast` shape. | | quantized (gather_qmv / gather_qmm — gather variants) | ✓ | ✓ | ✓ | `ffai/moe.rs` → `mt_moe_gather_qmm_int4` (int4 affine grouped-gather) + `mt_moe_gather_qmm_b{3,5,6,8}` (all bit-widths, scalar). **int4 perf**: `mt_moe_gather_qmm_mma_int4{,_bm16}` + `_m8` (decode) + `_m{16,32}` (PR #157 short-prefill, hand-unrolled `acc0..accN` cells — the DSL doesn't lower runtime-indexed mutable arrays), MPP scale-ups `bm{8,16,64}_mpp` (`ffai/moe_mpp{,_bm8,_bm64}.rs`). **int8 perf** (PR #154): pack-aligned `mt_moe_gather_qmm_mma_int8` (1-SG MMA decode) + `_bm16_mpp` + `_bm8_mpp` (direct-input cooperative tensors, M=8 forbids coop-tensor) + `_bm64_mpp` (4-SG 2×2 long-context prefill). All MPP kernels stage bf16 through `half` cooperative tensors via `coop_stage(T)`. Bare-tensor `kernels/ops/gather.rs` exists but is non-quantized. **Expert-indexed dequant GEMV** (PR #160): `dequant_gemv_int4_expert_indexed` — per-output-row expert selection for the gate/up FFN dispatch shape. | @@ -142,7 +142,7 @@ Local verification of NAX kernels is the developer's responsibility on M4+ hardw | sdpa_decode + learned attention sink (GPT-OSS-20B) | ✗ | ~ | ✓ | `ffai/sdpa_decode.rs` `has_sink` / `sink_logit` constexprs. GPT-OSS-20B's per-head learned attention-sink logit folds into the cross-simdgroup softmax denominator on-GPU as a virtual key — removing the host-side post-hoc rescale that previously cost a CPU sync per attention layer. | | gated_rmsnorm (fp32-in gated RMSNorm → activation dtype) | ✗ | ✗ | ✓ | `kernels/norm/gated_rmsnorm.rs` → `mt_gated_rmsnorm`. Fused Qwen3.5 / 3.6 GDN post-step `out = w·rmsNorm(y)·silu(z)`; `y` arrives fp32 (the `gated_delta` recurrence output). Closes the per-GDN-layer host-side CPU sync (~75 % of Qwen3.5/3.6 layers). | | conv2d (vision patch conv — im2col + tiled GEMM) | ✓ | ✓ | ✓ | `ffai/conv2d.rs` → `conv2d_patch14` / `conv2d_patch16` + `conv2d_generic`. NCHW input, OIHW weight; direct conv (implicit im2col, one thread per output). VLM front-end. | -| patch_embed (fused image unfold + linear projection) | ✗ | ✗ | ✓ | `ffai/patch_embed.rs` → `patch_embed`. Fused image-unfold + linear projection — gathers each patch's pixels and dots them with one weight row, no intermediate unfolded buffer. **MMA-tiled perf path** (PR #157): `ffai/patch_embed_mma.rs` → `patch_embed_mma` — implicit-patch-unfold + 4-SG 2×2 simdgroup-matrix MMA (`hidden` and `num_patches` divisible by 32); targets ViT-L/H shapes. | +| patch_embed (fused image unfold + linear projection) | ✗ | ✗ | ✓ | `kernels/gemm/patch_embed.rs` → `mt_patch_embed`. Fused image-unfold + linear projection — gathers each patch's pixels and dots them with one weight row, no intermediate unfolded buffer. **MMA-tiled perf path** (PR #157): `kernels/gemm/patch_embed_mma.rs` → `mt_patch_embed_mma` — implicit-patch-unfold + 4-SG 2×2 simdgroup-matrix MMA (`hidden` and `num_patches` divisible by 32); targets ViT-L/H shapes. | | rope_2d (2D positional RoPE for vision tokens) | ✓ | ✓ | ✓ | `kernels/rope/rope_2d.rs` → `mt_rope_2d`. 2D RoPE over a (row, col) token grid; head_dim split into row half + column half, each running rotate-half RoPE. VLM front-end. | | mel_spectrogram (STFT + log-Mel filterbank) | ✓ | ✓ | ✓ | `ffai/mel_spectrogram.rs` → `mel_spectrogram` (single-dispatch direct-DFT) + radix-FFT path `mel_stft_window` → `mt_fft_n{n_fft}` → `mel_filterbank` (three kernels, O(N log N)). Generic over `T` per PR #152. STT front-end. All four kernels are bounds-guarded (`idx < n_out`) for threadgroup-rounded dispatch. **Correctness of the two direct-DFT kernels (`mel_spectrogram`, `mel_spectrogram_magnitude`) is gated at f32**: their in-thread DFT hits spectrum cancellation nulls where the GPU's approximate `sin`/`cos` diverge from libm by orders of magnitude relative to the (near-zero) true power, which low-precision *input* rounding moves onto the null — flaky O(6–16) log error on a correct kernel. The kernels stay generic over `T`; f16/bf16 are covered by `mel_filterbank` (post-FFT, no in-thread cancellation). | | audio_conv1d (wide-stride 1D conv — STT patch embed) | ✓ | ✓ | ✓ | `ffai/audio_conv1d.rs` → `audio_conv1d`. Dense wide-stride multi-channel 1D conv (NCL); distinct from depthwise `conv1d_causal_step`. STT front-end. | diff --git a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md index 8f17ebe9..873d3b91 100644 --- a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md +++ b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md @@ -46,7 +46,9 @@ crates/metaltile-std/src/kernels/ ops/ ✅ DONE — elementwise/core primitives: binary · unary · ternary · copy · arange · random · reduce · arg_reduce · scan · indexing · gather/scatter · hadamard · fence · clamp · logsumexp · vector_add · axpy · strided · gated_activation - gemm/ gemm · gemv(_masked) · batched-projection (qkv / 4) · patch_embed · steel/gemm + gemm/ 🔨 DENSE DONE — gemm · gemv(_masked,_axpy) · patch_embed(_mma) · steel/gemm; + quantized matmuls (qmm/qgemv q4/q8/block_scaled/fp_quantized + batched qkv/4) + land in a follow-up pass (same folder; format-axis fold deferred, §7) sdpa/ ALL attention: bidirectional(+relpos/windowed/conformer) · decode(+d64..d512/ 2pass/batched/sink) · multi(+d256/tree-mask) · prefill_mma · flash_quantized · aura_flash · steel/attn @@ -157,7 +159,7 @@ payoff last: | Wave | Families | Rationale | Payoff | |---|---|---|---| | ✅ done | `convolution/`, `rope/`, `norm/`, `sampling/`, `ops/` | exemplar + all of wave 1 | 24k → ~1.6k | -| 2 | `gemm/`, `ssm/`, `audio/`, `vision/`, `kv_cache/` | moderate size, few cross-deps | medium | +| 2 | `gemm/` (🔨 dense in; quantized matmuls next), `ssm/`, `audio/`, `vision/`, `kv_cache/` | moderate size, few cross-deps | medium | | 3 | `sdpa/`, `moe/`, **`quant/`** | hardest axes (head-dim d64..d512; bm8/bm64×int8; the 30-format matrix) — most of the ~150k LOC | the bulk | ## 7. The `quant/` umbrella — collapsing the op × format matrix