From 8f5b2608c2fed666b0bf2594ae092105155e3770 Mon Sep 17 00:00:00 2001 From: Eric Kryski <599019+ekryski@users.noreply.github.com> Date: Sat, 13 Jun 2026 13:18:40 -0600 Subject: [PATCH 1/7] refactor(moe): migrate the moe family into kernels/moe/ MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Move all 35 moe-family files from ffai/ + mlx/ into kernels/moe/ (which #24 seeded with gather_q4 + sigmoid_bias): - orchestration (ex moe.rs): router_topk + permute/unpermute + 10 gather_qmm - routers: router_topk_biased (ex dsv4_router_topk), router_sigmoid_bias, router_sqrtsoftplus, sigmoid_bias - mpp grouped BGEMM: mpp(+int8/bm8/bm64/×int8/×block_scaled) + mpp_shared - gguf-format expert matmul: bgemm_{q2k,iq2xxs,q4}_*, gemv_{rows,ws}_*, gather_* - down combine: down_swiglu_accum, down_weighted_sum_f16 - expert-indexed + block-scaled: dequant_gemv_expert_indexed(_block_scaled), block_scaled (ex mlx/block_scaled_moe) Filenames drop the redundant moe_ prefix (folder provides it); kernel names keep mt_moe_*. Model-name purge: mt_dsv4_router_topk -> mt_moe_router_topk_biased (distinct from the generic mt_moe_router_topk: selects by the biased score, weights by the unbiased). Bare dequant_gemv_int4_expert_indexed -> mt_ prefix. Fixed mpp_shared intra-imports and ~27 consumer test files (grouped/mixed use-blocks included). Format-axis fold (§7) deferred. orchestration.rs (~4k lines) moves whole here; split follows next. --- crates/metaltile-std/src/ffai/mod.rs | 38 +---------- .../moe/bgemm_iq2xxs_bm64.rs} | 6 +- .../moe/bgemm_iq2xxs_mpp.rs} | 6 +- .../moe/bgemm_iq2xxs_view.rs} | 6 +- .../moe/bgemm_iq2xxs_view_u16_bm64.rs} | 6 +- .../moe/bgemm_q2k_bm64.rs} | 6 +- .../moe/bgemm_q2k_mpp.rs} | 6 +- .../moe/bgemm_q2k_view.rs} | 6 +- .../moe/bgemm_q2k_view_u16_bm64.rs} | 6 +- .../moe/bgemm_q4_bm64.rs} | 8 +-- .../moe/block_scaled.rs} | 0 .../moe}/dequant_gemv_expert_indexed.rs | 14 ++--- ...equant_gemv_expert_indexed_block_scaled.rs | 0 .../moe/down_swiglu_accum.rs} | 20 +++--- .../moe/down_weighted_sum_f16.rs} | 10 +-- .../moe/gather_down_q2k.rs} | 6 +- .../moe/gather_gemv_iq2xxs.rs} | 10 +-- .../moe/gemv_rows_iq2xxs.rs} | 10 +-- .../moe/gemv_rows_q2k.rs} | 8 +-- .../moe/gemv_rows_view_iq2xxs.rs} | 12 ++-- .../moe/gemv_ws_iq2xxs.rs} | 6 +- .../moe/gemv_ws_q2k.rs} | 8 +-- crates/metaltile-std/src/kernels/moe/mod.rs | 63 +++++++++++++++++-- .../{ffai/moe_mpp.rs => kernels/moe/mpp.rs} | 4 +- .../moe/mpp_block_scaled.rs} | 0 .../moe/mpp_bm64.rs} | 4 +- .../moe/mpp_bm64_block_scaled.rs} | 0 .../moe/mpp_bm64_int8.rs} | 4 +- .../moe_mpp_bm8.rs => kernels/moe/mpp_bm8.rs} | 4 +- .../moe/mpp_bm8_block_scaled.rs} | 0 .../moe/mpp_bm8_int8.rs} | 4 +- .../moe/mpp_int8.rs} | 4 +- .../moe/mpp_shared.rs} | 0 .../moe.rs => kernels/moe/orchestration.rs} | 0 .../moe/router_sigmoid_bias.rs} | 10 +-- .../moe/router_sqrtsoftplus.rs} | 12 ++-- .../moe/router_topk_biased.rs} | 8 +-- .../src/kernels/moe/sigmoid_bias.rs | 2 +- crates/metaltile-std/src/mlx/mod.rs | 5 +- .../tests/bm64_vs_gemvrows_iq2xxs.rs | 12 ++-- .../tests/dsv4_router_topk_correctness.rs | 6 +- .../moe_bgemm_iq2xxs_bm64_correctness.rs | 12 ++-- .../tests/moe_bgemm_iq2xxs_mpp_correctness.rs | 6 +- .../moe_bgemm_iq2xxs_view_correctness.rs | 8 +-- .../tests/moe_bgemm_q2k_bm64_correctness.rs | 12 ++-- .../tests/moe_bgemm_q2k_mpp_correctness.rs | 6 +- .../tests/moe_bgemm_q2k_view_correctness.rs | 14 ++--- .../tests/moe_bm64_ragged_correctness.rs | 10 +-- .../tests/moe_gather_down_q2k_correctness.rs | 6 +- .../moe_gather_gemv_iq2xxs_correctness.rs | 6 +- ...moe_gather_qmm_int4_m16_m32_correctness.rs | 2 +- .../tests/moe_gather_qmm_microbench.rs | 2 +- ...moe_gather_qmm_mma_bitwidth_correctness.rs | 4 +- .../moe_gather_qmm_mpp_bm64_correctness.rs | 10 +-- ...oe_gather_qmm_mpp_bm64_int8_correctness.rs | 10 +-- .../moe_gather_qmm_mpp_bm8_correctness.rs | 6 +- ...moe_gather_qmm_mpp_bm8_int8_correctness.rs | 6 +- .../tests/moe_gather_qmm_mpp_correctness.rs | 8 +-- .../moe_gather_qmm_mpp_int8_correctness.rs | 6 +- .../tests/moe_gemv_rows_iq2xxs_correctness.rs | 6 +- .../tests/moe_gemv_rows_q2k_correctness.rs | 14 ++--- .../moe_gemv_rows_view_iq2xxs_correctness.rs | 14 ++--- .../tests/moe_gemv_ws_iq2xxs_correctness.rs | 6 +- .../tests/moe_gemv_ws_q2k_correctness.rs | 14 ++--- .../tests/moe_q2k_view_u16_correctness.rs | 10 +-- .../tests/moe_view_u16_correctness.rs | 10 +-- docs/specs/KERNEL_AUDIT.md | 12 ++-- docs/specs/KERNEL_CONSOLIDATION_PLAN.md | 11 ++-- 68 files changed, 297 insertions(+), 274 deletions(-) rename crates/metaltile-std/src/{ffai/moe_bgemm_iq2xxs_bm64.rs => kernels/moe/bgemm_iq2xxs_bm64.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_bgemm_iq2xxs_mpp.rs => kernels/moe/bgemm_iq2xxs_mpp.rs} (97%) rename crates/metaltile-std/src/{ffai/moe_bgemm_iq2xxs_view.rs => kernels/moe/bgemm_iq2xxs_view.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_bgemm_iq2xxs_view_u16_bm64.rs => kernels/moe/bgemm_iq2xxs_view_u16_bm64.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_bgemm_q2k_bm64.rs => kernels/moe/bgemm_q2k_bm64.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_bgemm_q2k_mpp.rs => kernels/moe/bgemm_q2k_mpp.rs} (97%) rename crates/metaltile-std/src/{ffai/moe_bgemm_q2k_view.rs => kernels/moe/bgemm_q2k_view.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_bgemm_q2k_view_u16_bm64.rs => kernels/moe/bgemm_q2k_view_u16_bm64.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_bgemm_q4_bm64.rs => kernels/moe/bgemm_q4_bm64.rs} (98%) rename crates/metaltile-std/src/{mlx/block_scaled_moe.rs => kernels/moe/block_scaled.rs} (100%) rename crates/metaltile-std/src/{ffai => kernels/moe}/dequant_gemv_expert_indexed.rs (95%) rename crates/metaltile-std/src/{ffai => kernels/moe}/dequant_gemv_expert_indexed_block_scaled.rs (100%) rename crates/metaltile-std/src/{ffai/moe_down_swiglu_accum.rs => kernels/moe/down_swiglu_accum.rs} (97%) rename crates/metaltile-std/src/{ffai/moe_down_weighted_sum_f16.rs => kernels/moe/down_weighted_sum_f16.rs} (97%) rename crates/metaltile-std/src/{ffai/moe_gather_down_q2k.rs => kernels/moe/gather_down_q2k.rs} (97%) rename crates/metaltile-std/src/{ffai/moe_gather_gemv_iq2xxs.rs => kernels/moe/gather_gemv_iq2xxs.rs} (96%) rename crates/metaltile-std/src/{ffai/moe_gemv_rows_iq2xxs.rs => kernels/moe/gemv_rows_iq2xxs.rs} (93%) rename crates/metaltile-std/src/{ffai/moe_gemv_rows_q2k.rs => kernels/moe/gemv_rows_q2k.rs} (95%) rename crates/metaltile-std/src/{ffai/moe_gemv_rows_view_iq2xxs.rs => kernels/moe/gemv_rows_view_iq2xxs.rs} (96%) rename crates/metaltile-std/src/{ffai/moe_gemv_ws_iq2xxs.rs => kernels/moe/gemv_ws_iq2xxs.rs} (97%) rename crates/metaltile-std/src/{ffai/moe_gemv_ws_q2k.rs => kernels/moe/gemv_ws_q2k.rs} (96%) rename crates/metaltile-std/src/{ffai/moe_mpp.rs => kernels/moe/mpp.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_mpp_block_scaled.rs => kernels/moe/mpp_block_scaled.rs} (100%) rename crates/metaltile-std/src/{ffai/moe_mpp_bm64.rs => kernels/moe/mpp_bm64.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_mpp_bm64_block_scaled.rs => kernels/moe/mpp_bm64_block_scaled.rs} (100%) rename crates/metaltile-std/src/{ffai/moe_mpp_bm64_int8.rs => kernels/moe/mpp_bm64_int8.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_mpp_bm8.rs => kernels/moe/mpp_bm8.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_mpp_bm8_block_scaled.rs => kernels/moe/mpp_bm8_block_scaled.rs} (100%) rename crates/metaltile-std/src/{ffai/moe_mpp_bm8_int8.rs => kernels/moe/mpp_bm8_int8.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_mpp_int8.rs => kernels/moe/mpp_int8.rs} (98%) rename crates/metaltile-std/src/{ffai/moe_mpp_shared.rs => kernels/moe/mpp_shared.rs} (100%) rename crates/metaltile-std/src/{ffai/moe.rs => kernels/moe/orchestration.rs} (100%) rename crates/metaltile-std/src/{ffai/moe_router_sigmoid_bias.rs => kernels/moe/router_sigmoid_bias.rs} (94%) rename crates/metaltile-std/src/{ffai/moe_router_sqrtsoftplus.rs => kernels/moe/router_sqrtsoftplus.rs} (95%) rename crates/metaltile-std/src/{ffai/dsv4_router_topk.rs => kernels/moe/router_topk_biased.rs} (95%) diff --git a/crates/metaltile-std/src/ffai/mod.rs b/crates/metaltile-std/src/ffai/mod.rs index 7d1c8ea6..ecb1ca01 100644 --- a/crates/metaltile-std/src/ffai/mod.rs +++ b/crates/metaltile-std/src/ffai/mod.rs @@ -27,8 +27,6 @@ pub mod aura_value; // batched_* projection GEMV/GEMM + dequant_gemv migrated to kernels/gemm/. pub mod dequant_gather; pub mod dequant_gather_block_scaled; -pub mod dequant_gemv_expert_indexed; -pub mod dequant_gemv_expert_indexed_block_scaled; pub mod dsv4_compressor_pool; pub mod dsv4_csa_sdpa_decode; pub mod dsv4_fp8_block_dequant; @@ -37,7 +35,6 @@ pub mod dsv4_indexer_topk; pub mod dsv4_mhc; pub mod dsv4_mhc_sinkhorn_split; pub mod dsv4_mxfp4_dequant; -pub mod dsv4_router_topk; pub mod dsv4_swiglu_limit; pub mod ffai_dequant_q4; pub mod flash_block_scaled_sdpa; @@ -51,38 +48,9 @@ pub mod gguf_dequant_q2_k; pub mod gguf_dequant_q8_0; pub mod gguf_iq2_xxs_extract_qs; pub mod leaky_relu; -pub mod moe; -pub mod moe_bgemm_iq2xxs_bm64; -pub mod moe_bgemm_iq2xxs_mpp; -pub mod moe_bgemm_iq2xxs_view; -pub mod moe_bgemm_iq2xxs_view_u16_bm64; -pub mod moe_bgemm_q2k_bm64; -pub mod moe_bgemm_q2k_mpp; -pub mod moe_bgemm_q2k_view; -pub mod moe_bgemm_q2k_view_u16_bm64; -pub mod moe_bgemm_q4_bm64; -pub mod moe_down_swiglu_accum; -pub mod moe_down_weighted_sum_f16; -pub mod moe_gather_down_q2k; -pub mod moe_gather_gemv_iq2xxs; -pub mod moe_gemv_rows_iq2xxs; -pub mod moe_gemv_rows_q2k; -pub mod moe_gemv_rows_view_iq2xxs; -pub mod moe_gemv_ws_iq2xxs; -pub mod moe_gemv_ws_q2k; -pub mod moe_mpp; -pub mod moe_mpp_block_scaled; -pub mod moe_mpp_bm64; -pub mod moe_mpp_bm64_block_scaled; -pub mod moe_mpp_bm64_int8; -pub mod moe_mpp_bm8; -pub mod moe_mpp_bm8_block_scaled; -pub mod moe_mpp_bm8_int8; -pub mod moe_mpp_int8; -pub mod moe_mpp_shared; -pub mod moe_router_sigmoid_bias; -pub mod moe_router_sqrtsoftplus; -// patch_embed_block_scaled / patch_embed_mma_block_scaled → kernels/gemm/. +// moe family (moe* · dsv4_router_topk · dequant_gemv_expert_indexed*) → +// kernels/moe/. patch_embed_block_scaled / patch_embed_mma_block_scaled → +// kernels/gemm/. pub mod sdpa_bidirectional; pub mod sdpa_bidirectional_d128_relpos; pub mod sdpa_bidirectional_windowed; diff --git a/crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_bm64.rs b/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_bm64.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_bm64.rs rename to crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_bm64.rs index 42f91778..7a3eecfa 100644 --- a/crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_bm64.rs +++ b/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_bm64.rs @@ -16,7 +16,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_bgemm_iq2xxs_bm64( +pub fn mt_moe_bgemm_iq2xxs_bm64( x: Tensor, qs: Tensor, d_f32: Tensor, @@ -163,7 +163,7 @@ pub fn ffai_moe_bgemm_iq2xxs_bm64( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_bgemm_iq2xxs_bm64; + use super::mt_moe_bgemm_iq2xxs_bm64; #[bench(dtypes = [f32, f16, bf16])] fn bench_bgemm_iq2xxs_bm64(dt: DType) -> BenchSetup { @@ -172,7 +172,7 @@ pub mod kernel_benches { let n_out = 2048usize; let t_rows = 256usize; let nblk = n_out * k_in / 256; - BenchSetup::new(ffai_moe_bgemm_iq2xxs_bm64::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_bgemm_iq2xxs_bm64::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", t_rows * k_in, dt)) .buffer(BenchBuffer::random("qs", n_experts * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_mpp.rs b/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_mpp.rs similarity index 97% rename from crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_mpp.rs rename to crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_mpp.rs index 7355b0b1..c6ceebf4 100644 --- a/crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_mpp.rs +++ b/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_mpp.rs @@ -22,7 +22,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_gather_bgemm_iq2xxs_mpp( +pub fn mt_moe_gather_bgemm_iq2xxs_mpp( x: Tensor, qs: Tensor, d_f32: Tensor, @@ -154,7 +154,7 @@ pub fn ffai_moe_gather_bgemm_iq2xxs_mpp( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_gather_bgemm_iq2xxs_mpp; + use super::mt_moe_gather_bgemm_iq2xxs_mpp; #[bench(dtypes = [f32, f16, bf16])] fn bench_bgemm_iq2xxs_mpp(dt: DType) -> BenchSetup { @@ -163,7 +163,7 @@ pub mod kernel_benches { let n_out = 2048usize; let t_rows = 256usize; let nblk = n_out * k_in / 256; - BenchSetup::new(ffai_moe_gather_bgemm_iq2xxs_mpp::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gather_bgemm_iq2xxs_mpp::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", t_rows * k_in, dt)) .buffer(BenchBuffer::random("qs", n_experts * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_view.rs b/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_view.rs rename to crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view.rs index 8aa2d9d7..1489b41f 100644 --- a/crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_view.rs +++ b/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view.rs @@ -21,7 +21,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_bgemm_iq2xxs_view( +pub fn mt_moe_bgemm_iq2xxs_view( x: Tensor, view_u8: Tensor, grid: Tensor, @@ -174,7 +174,7 @@ pub fn ffai_moe_bgemm_iq2xxs_view( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_bgemm_iq2xxs_view; + use super::mt_moe_bgemm_iq2xxs_view; #[bench(dtypes = [f32, f16, bf16])] fn bench_bgemm_iq2xxs_view(dt: DType) -> BenchSetup { @@ -185,7 +185,7 @@ pub mod kernel_benches { let nblk = n_out * k_in / 256; // view holds n_experts × nblk IQ2 blocks of 66 bytes each. let view_bytes = n_experts * nblk * 66; - BenchSetup::new(ffai_moe_bgemm_iq2xxs_view::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_bgemm_iq2xxs_view::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", t_rows * k_in, dt)) .buffer(BenchBuffer::random("view_u8", view_bytes, DType::U8)) diff --git a/crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_view_u16_bm64.rs b/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view_u16_bm64.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_view_u16_bm64.rs rename to crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view_u16_bm64.rs index 15624468..0552f8e1 100644 --- a/crates/metaltile-std/src/ffai/moe_bgemm_iq2xxs_view_u16_bm64.rs +++ b/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view_u16_bm64.rs @@ -16,7 +16,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_bgemm_iq2xxs_view_u16_bm64( +pub fn mt_moe_bgemm_iq2xxs_view_u16_bm64( x: Tensor, view_u16: Tensor, view_f16: Tensor, @@ -165,7 +165,7 @@ pub fn ffai_moe_bgemm_iq2xxs_view_u16_bm64( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_bgemm_iq2xxs_view_u16_bm64; + use super::mt_moe_bgemm_iq2xxs_view_u16_bm64; #[bench(dtypes = [f32, f16, bf16])] fn bench(dt: DType) -> BenchSetup { @@ -174,7 +174,7 @@ pub mod kernel_benches { let n_out = 2048usize; let t_rows = 256usize; let nblk = n_out * k_in / 256; - BenchSetup::new(ffai_moe_bgemm_iq2xxs_view_u16_bm64::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_bgemm_iq2xxs_view_u16_bm64::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", t_rows * k_in, dt)) .buffer(BenchBuffer::random("view_u16", n_experts * nblk * 33, DType::U16)) diff --git a/crates/metaltile-std/src/ffai/moe_bgemm_q2k_bm64.rs b/crates/metaltile-std/src/kernels/moe/bgemm_q2k_bm64.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_bgemm_q2k_bm64.rs rename to crates/metaltile-std/src/kernels/moe/bgemm_q2k_bm64.rs index 466fa952..f342c1e0 100644 --- a/crates/metaltile-std/src/ffai/moe_bgemm_q2k_bm64.rs +++ b/crates/metaltile-std/src/kernels/moe/bgemm_q2k_bm64.rs @@ -10,7 +10,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_bgemm_q2k_bm64( +pub fn mt_moe_bgemm_q2k_bm64( x: Tensor, qs: Tensor, scales: Tensor, @@ -160,7 +160,7 @@ pub fn ffai_moe_bgemm_q2k_bm64( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_bgemm_q2k_bm64; + use super::mt_moe_bgemm_q2k_bm64; #[bench(dtypes = [f32, f16, bf16])] fn bench_bgemm_q2k_bm64(dt: DType) -> BenchSetup { @@ -169,7 +169,7 @@ pub mod kernel_benches { let n_out = 4096usize; let t_rows = 256usize; let nblk = n_out * k_in / 256; - BenchSetup::new(ffai_moe_bgemm_q2k_bm64::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_bgemm_q2k_bm64::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", t_rows * k_in, dt)) .buffer(BenchBuffer::random("qs", n_experts * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/ffai/moe_bgemm_q2k_mpp.rs b/crates/metaltile-std/src/kernels/moe/bgemm_q2k_mpp.rs similarity index 97% rename from crates/metaltile-std/src/ffai/moe_bgemm_q2k_mpp.rs rename to crates/metaltile-std/src/kernels/moe/bgemm_q2k_mpp.rs index 08da05e7..159a0960 100644 --- a/crates/metaltile-std/src/ffai/moe_bgemm_q2k_mpp.rs +++ b/crates/metaltile-std/src/kernels/moe/bgemm_q2k_mpp.rs @@ -16,7 +16,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_gather_bgemm_q2k_mpp( +pub fn mt_moe_gather_bgemm_q2k_mpp( x: Tensor, qs: Tensor, scales: Tensor, @@ -152,7 +152,7 @@ pub fn ffai_moe_gather_bgemm_q2k_mpp( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_gather_bgemm_q2k_mpp; + use super::mt_moe_gather_bgemm_q2k_mpp; #[bench(dtypes = [f32, f16, bf16])] fn bench_bgemm_q2k_mpp(dt: DType) -> BenchSetup { @@ -161,7 +161,7 @@ pub mod kernel_benches { let n_out = 4096usize; let t_rows = 256usize; let nblk = n_out * k_in / 256; - BenchSetup::new(ffai_moe_gather_bgemm_q2k_mpp::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gather_bgemm_q2k_mpp::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", t_rows * k_in, dt)) .buffer(BenchBuffer::random("qs", n_experts * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/ffai/moe_bgemm_q2k_view.rs b/crates/metaltile-std/src/kernels/moe/bgemm_q2k_view.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_bgemm_q2k_view.rs rename to crates/metaltile-std/src/kernels/moe/bgemm_q2k_view.rs index e166dce6..be0303e5 100644 --- a/crates/metaltile-std/src/ffai/moe_bgemm_q2k_view.rs +++ b/crates/metaltile-std/src/kernels/moe/bgemm_q2k_view.rs @@ -17,7 +17,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_bgemm_q2k_view( +pub fn mt_moe_bgemm_q2k_view( x: Tensor, view_u8: Tensor, indices: Tensor, @@ -167,7 +167,7 @@ pub fn ffai_moe_bgemm_q2k_view( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_bgemm_q2k_view; + use super::mt_moe_bgemm_q2k_view; #[bench(dtypes = [f32, f16, bf16])] fn bench_bgemm_q2k_view(dt: DType) -> BenchSetup { @@ -177,7 +177,7 @@ pub mod kernel_benches { let t_rows = 256usize; let nblk = n_out * k_in / 256; let view_bytes = n_experts * nblk * 84; - BenchSetup::new(ffai_moe_bgemm_q2k_view::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_bgemm_q2k_view::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", t_rows * k_in, dt)) .buffer(BenchBuffer::random("view_u8", view_bytes, DType::U8)) diff --git a/crates/metaltile-std/src/ffai/moe_bgemm_q2k_view_u16_bm64.rs b/crates/metaltile-std/src/kernels/moe/bgemm_q2k_view_u16_bm64.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_bgemm_q2k_view_u16_bm64.rs rename to crates/metaltile-std/src/kernels/moe/bgemm_q2k_view_u16_bm64.rs index 1c6b2c85..ac365a53 100644 --- a/crates/metaltile-std/src/ffai/moe_bgemm_q2k_view_u16_bm64.rs +++ b/crates/metaltile-std/src/kernels/moe/bgemm_q2k_view_u16_bm64.rs @@ -16,7 +16,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_bgemm_q2k_view_u16_bm64( +pub fn mt_moe_bgemm_q2k_view_u16_bm64( x: Tensor, view_u16: Tensor, view_f16: Tensor, @@ -168,7 +168,7 @@ pub fn ffai_moe_bgemm_q2k_view_u16_bm64( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_bgemm_q2k_view_u16_bm64; + use super::mt_moe_bgemm_q2k_view_u16_bm64; #[bench(dtypes = [f32, f16, bf16])] fn bench(dt: DType) -> BenchSetup { @@ -177,7 +177,7 @@ pub mod kernel_benches { let n_out = 4096usize; let t_rows = 256usize; let nblk = n_out * k_in / 256; - BenchSetup::new(ffai_moe_bgemm_q2k_view_u16_bm64::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_bgemm_q2k_view_u16_bm64::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", t_rows * k_in, dt)) .buffer(BenchBuffer::random("view_u16", n_experts * nblk * 42, DType::U16)) diff --git a/crates/metaltile-std/src/ffai/moe_bgemm_q4_bm64.rs b/crates/metaltile-std/src/kernels/moe/bgemm_q4_bm64.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_bgemm_q4_bm64.rs rename to crates/metaltile-std/src/kernels/moe/bgemm_q4_bm64.rs index ff12555c..d57a1ef7 100644 --- a/crates/metaltile-std/src/ffai/moe_bgemm_q4_bm64.rs +++ b/crates/metaltile-std/src/kernels/moe/bgemm_q4_bm64.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 //! HIGH-THROUGHPUT amortized MoE Q4 grouped BGEMM — bm64 tiling (64×64×32, //! 4 simdgroups) with the bench's signed-4-bit dequant. The Q4 twin of -//! `ffai_moe_bgemm_q2k_bm64`: processes a 64-row M-tile whose rows are +//! `mt_moe_bgemm_q2k_bm64`: processes a 64-row M-tile whose rows are //! PRE-SORTED by expert id (`indices[row]`), finds contiguous same-expert //! sub-runs, and runs one MMA GEMM per sub-run against that expert's weights. //! This replaces the per-token MoE gather loop (which was 72% of prefill time @@ -24,7 +24,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_bgemm_q4_bm64( +pub fn mt_moe_bgemm_q4_bm64( x: Tensor, qs: Tensor, scales: Tensor, @@ -159,7 +159,7 @@ pub fn ffai_moe_bgemm_q4_bm64( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::ffai_moe_bgemm_q4_bm64; + use super::mt_moe_bgemm_q4_bm64; use crate::utils::pack_f32; fn quantize_q4(w: &[f32], m: usize, k: usize) -> (Vec, Vec) { @@ -231,7 +231,7 @@ pub mod kernel_tests { let qs_bytes: Vec = qs.iter().flat_map(|x| x.to_le_bytes()).collect(); let idx_bytes: Vec = idx.iter().flat_map(|x| x.to_le_bytes()).collect(); let _ = bpr; - TestSetup::new(ffai_moe_bgemm_q4_bm64::kernel_ir_for(dt)) + TestSetup::new(mt_moe_bgemm_q4_bm64::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .input(TestBuffer::from_vec("x", pack_f32(&xv, dt), dt)) .input(TestBuffer::from_vec("qs", qs_bytes, DType::U32)) diff --git a/crates/metaltile-std/src/mlx/block_scaled_moe.rs b/crates/metaltile-std/src/kernels/moe/block_scaled.rs similarity index 100% rename from crates/metaltile-std/src/mlx/block_scaled_moe.rs rename to crates/metaltile-std/src/kernels/moe/block_scaled.rs diff --git a/crates/metaltile-std/src/ffai/dequant_gemv_expert_indexed.rs b/crates/metaltile-std/src/kernels/moe/dequant_gemv_expert_indexed.rs similarity index 95% rename from crates/metaltile-std/src/ffai/dequant_gemv_expert_indexed.rs rename to crates/metaltile-std/src/kernels/moe/dequant_gemv_expert_indexed.rs index 45f78caf..8f27de7b 100644 --- a/crates/metaltile-std/src/ffai/dequant_gemv_expert_indexed.rs +++ b/crates/metaltile-std/src/kernels/moe/dequant_gemv_expert_indexed.rs @@ -55,7 +55,7 @@ use metaltile::kernel; #[kernel] -pub fn dequant_gemv_int4_expert_indexed( +pub fn mt_dequant_gemv_int4_expert_indexed( weights_stacked: Tensor, scales_stacked: Tensor, biases_stacked: Tensor, @@ -104,7 +104,7 @@ pub fn dequant_gemv_int4_expert_indexed( } } -/// New-syntax correctness test for `dequant_gemv_int4_expert_indexed` — the +/// New-syntax correctness test for `mt_dequant_gemv_int4_expert_indexed` — the /// per-expert-indexed int4 dequant GEMV. Reduction-mode (one threadgroup per /// output row, `reduce_sum` across the threadgroup). /// @@ -118,7 +118,7 @@ pub fn dequant_gemv_int4_expert_indexed( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::dequant_gemv_int4_expert_indexed; + use super::mt_dequant_gemv_int4_expert_indexed; use crate::utils::{pack_f32, unpack_f32}; fn u32_bytes(v: &[u32]) -> Vec { v.iter().flat_map(|x| x.to_le_bytes()).collect() } @@ -198,7 +198,7 @@ pub mod kernel_tests { let b = unpack_f32(&pack_f32(&biases_f, dt), dt); let x = unpack_f32(&pack_f32(&input_f, dt), dt); let expected = oracle(&w, &s, &b, &x, expert, out_dim, in_dim, group_size); - TestSetup::new(dequant_gemv_int4_expert_indexed::kernel_ir_for(dt)) + TestSetup::new(mt_dequant_gemv_int4_expert_indexed::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .input(TestBuffer::from_vec("weights_stacked", u32_bytes(&w), DType::U32)) .input(TestBuffer::from_vec("scales_stacked", pack_f32(&scales_f, dt), dt)) @@ -214,12 +214,12 @@ pub mod kernel_tests { } } -/// New-syntax benchmark for `dequant_gemv_int4_expert_indexed`. Production-ish +/// New-syntax benchmark for `mt_dequant_gemv_int4_expert_indexed`. Production-ish /// shape (out_dim/in_dim 4096, group_size 64, 8 experts). One TG per output row. pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::dequant_gemv_int4_expert_indexed; + use super::mt_dequant_gemv_int4_expert_indexed; #[bench(dtypes = [f32, f16, bf16])] fn bench_dequant_gemv_int4_expert_indexed(dt: DType) -> BenchSetup { @@ -230,7 +230,7 @@ pub mod kernel_benches { // Active stream: one expert's weight slab + its scales/biases + input + output. let bytes = out_dim * packs_per_row * 4 + 2 * out_dim * n_groups * sz + in_dim * sz + out_dim * sz; - BenchSetup::new(dequant_gemv_int4_expert_indexed::kernel_ir_for(dt)) + BenchSetup::new(mt_dequant_gemv_int4_expert_indexed::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random( "weights_stacked", diff --git a/crates/metaltile-std/src/ffai/dequant_gemv_expert_indexed_block_scaled.rs b/crates/metaltile-std/src/kernels/moe/dequant_gemv_expert_indexed_block_scaled.rs similarity index 100% rename from crates/metaltile-std/src/ffai/dequant_gemv_expert_indexed_block_scaled.rs rename to crates/metaltile-std/src/kernels/moe/dequant_gemv_expert_indexed_block_scaled.rs diff --git a/crates/metaltile-std/src/ffai/moe_down_swiglu_accum.rs b/crates/metaltile-std/src/kernels/moe/down_swiglu_accum.rs similarity index 97% rename from crates/metaltile-std/src/ffai/moe_down_swiglu_accum.rs rename to crates/metaltile-std/src/kernels/moe/down_swiglu_accum.rs index d814040e..fa55e18e 100644 --- a/crates/metaltile-std/src/ffai/moe_down_swiglu_accum.rs +++ b/crates/metaltile-std/src/kernels/moe/down_swiglu_accum.rs @@ -7,7 +7,7 @@ //! as of ITER 80) into ONE kernel launch: //! //! 1. `mt_swiglu` (many=8): `inner[k][d] = silu(gate[k][d]) * up[k][d]` -//! 2. `ffai_dequant_gemv_int4_expert_indexed` (many=8): per slot k, +//! 2. `mt_dequant_gemv_int4_expert_indexed` (many=8): per slot k, //! `down_out[k] = W_down[expert[k]] · inner[k]` (out_dim = hidden) //! 3. `mt_scalar_fma_chain8`: //! `acc[i] = Σ_{k=0..8} scalar[k] * down_out[k][i]` @@ -75,7 +75,7 @@ //! floating-point reorder of the per-thread reduction) to: //! //! for k in 0..8: tmp_k = mt_swiglu(gate_k, up_k) -//! for k in 0..8: down_k = ffai_dequant_gemv_int4_expert_indexed( +//! for k in 0..8: down_k = mt_dequant_gemv_int4_expert_indexed( //! W, S, B, tmp_k, expert_indices[k:k+1]) //! out = mt_scalar_fma_chain8(slot_weights[0:1], down_0, ..., //! slot_weights[7:8], down_7) @@ -155,7 +155,7 @@ macro_rules! define_moe_down_swiglu_accum_chain8 { ) => { #[kernel] #[allow(clippy::too_many_arguments)] - pub fn ffai_moe_down_swiglu_accum_int4_chain8( + pub fn mt_moe_down_swiglu_accum_int4_chain8( gate_0: Tensor, up_0: Tensor, gate_1: Tensor, @@ -189,7 +189,7 @@ macro_rules! define_moe_down_swiglu_accum_chain8 { // accumulation precision. threadgroup_alloc("tg_inner", 768, "f32"); - // Int4 dequant constants, match `dequant_gemv_int4_expert_indexed`. + // Int4 dequant constants, match `mt_dequant_gemv_int4_expert_indexed`. let vals_per_pack = 8u32; let mask = 0xFu32; let row = program_id::<0>(); @@ -309,7 +309,7 @@ define_moe_down_swiglu_accum_chain8!( ); /// New-syntax correctness test for the fused MoE decode kernel -/// (`ffai_moe_down_swiglu_accum_int4_chain8`). The 8-way fusion has a clean +/// (`mt_moe_down_swiglu_accum_int4_chain8`). The 8-way fusion has a clean /// closed-form oracle: for each output row `i`, /// /// out[i] = Σ_{k=0..8} slot_weights[k] @@ -325,7 +325,7 @@ define_moe_down_swiglu_accum_chain8!( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::ffai_moe_down_swiglu_accum_int4_chain8; + use super::mt_moe_down_swiglu_accum_int4_chain8; use crate::utils::{pack_f32, unpack_f32}; /// Top-k slot count this kernel fuses (8-way chain). @@ -460,7 +460,7 @@ pub mod kernel_tests { group_size, ); - let mut su = TestSetup::new(ffai_moe_down_swiglu_accum_int4_chain8::kernel_ir_for(dt)) + let mut su = TestSetup::new(mt_moe_down_swiglu_accum_int4_chain8::kernel_ir_for(dt)) .mode(KernelMode::Reduction); for k in 0..N_SLOTS { su = su @@ -488,7 +488,7 @@ pub mod kernel_tests { } /// New-syntax benchmark for the fused MoE decode kernel -/// (`ffai_moe_down_swiglu_accum_int4_chain8`). Bench-only: the 8-way +/// (`mt_moe_down_swiglu_accum_int4_chain8`). Bench-only: the 8-way /// SwiGLU + indexed int4 down-projection + scalar-FMA-chain fusion has no /// clean single-stage oracle — its end-to-end correctness is validated in /// FFAI integration tests and by the in-source `#[test_kernel]`s against the @@ -502,7 +502,7 @@ pub mod kernel_tests { pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_down_swiglu_accum_int4_chain8; + use super::mt_moe_down_swiglu_accum_int4_chain8; /// Lanes per threadgroup — the caller-picked `lsize` (typically 128). const LSIZE: u32 = 128; @@ -527,7 +527,7 @@ pub mod kernel_benches { + 2 * n_experts * out_dim * n_groups * sz + out_dim * sz; - let mut bs = BenchSetup::new(ffai_moe_down_swiglu_accum_int4_chain8::kernel_ir_for(dt)) + let mut bs = BenchSetup::new(mt_moe_down_swiglu_accum_int4_chain8::kernel_ir_for(dt)) .mode(KernelMode::Reduction); // 8 gate/up activation pairs. for k in 0..N_SLOTS { diff --git a/crates/metaltile-std/src/ffai/moe_down_weighted_sum_f16.rs b/crates/metaltile-std/src/kernels/moe/down_weighted_sum_f16.rs similarity index 97% rename from crates/metaltile-std/src/ffai/moe_down_weighted_sum_f16.rs rename to crates/metaltile-std/src/kernels/moe/down_weighted_sum_f16.rs index 8ba5b631..eb836d30 100644 --- a/crates/metaltile-std/src/ffai/moe_down_weighted_sum_f16.rs +++ b/crates/metaltile-std/src/kernels/moe/down_weighted_sum_f16.rs @@ -27,7 +27,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_down_weighted_sum_6( +pub fn mt_moe_down_weighted_sum_6( down_0: Tensor, inner_0: Tensor, down_1: Tensor, @@ -271,7 +271,7 @@ pub fn ffai_moe_down_weighted_sum_6( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::ffai_moe_down_weighted_sum_6; + use super::mt_moe_down_weighted_sum_6; use crate::utils::{pack_f32, unpack_f32}; fn setup(m: usize, k: usize, dt: DType) -> TestSetup { @@ -300,7 +300,7 @@ pub mod kernel_tests { s }) .collect(); - TestSetup::new(ffai_moe_down_weighted_sum_6::kernel_ir_for(dt)) + TestSetup::new(mt_moe_down_weighted_sum_6::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .input(TestBuffer::from_vec("down_0", pack_f32(&dws[0], dt), dt)) .input(TestBuffer::from_vec("inner_0", pack_f32(&inns[0], dt), dt)) @@ -328,12 +328,12 @@ pub mod kernel_tests { pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_down_weighted_sum_6; + use super::mt_moe_down_weighted_sum_6; #[bench(dtypes = [f32, f16, bf16])] fn bench_mds(dt: DType) -> BenchSetup { let (m, k) = (4096usize, 2048usize); - BenchSetup::new(ffai_moe_down_weighted_sum_6::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_down_weighted_sum_6::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("down_0", m * k, dt)) .buffer(BenchBuffer::random("inner_0", k, dt)) diff --git a/crates/metaltile-std/src/ffai/moe_gather_down_q2k.rs b/crates/metaltile-std/src/kernels/moe/gather_down_q2k.rs similarity index 97% rename from crates/metaltile-std/src/ffai/moe_gather_down_q2k.rs rename to crates/metaltile-std/src/kernels/moe/gather_down_q2k.rs index 73fc9671..a469f6bc 100644 --- a/crates/metaltile-std/src/ffai/moe_gather_down_q2k.rs +++ b/crates/metaltile-std/src/kernels/moe/gather_down_q2k.rs @@ -35,7 +35,7 @@ use metaltile::kernel; #[kernel] -pub fn ffai_moe_gather_down_q2k( +pub fn mt_moe_gather_down_q2k( inners_all: Tensor, qs_all: Tensor, scales_all: Tensor, @@ -106,7 +106,7 @@ pub fn ffai_moe_gather_down_q2k( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_gather_down_q2k; + use super::mt_moe_gather_down_q2k; // n_slots=6; production down dims (m_out=4096, k_in=2048). #[bench(dtypes = [f32, f16, bf16])] @@ -115,7 +115,7 @@ pub mod kernel_benches { let m_out = 4096usize; let k_in = 2048usize; let nblk = m_out * (k_in / 256); - BenchSetup::new(ffai_moe_gather_down_q2k::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gather_down_q2k::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("inners_all", n_slots * k_in, dt)) .buffer(BenchBuffer::random("qs_all", n_slots * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/ffai/moe_gather_gemv_iq2xxs.rs b/crates/metaltile-std/src/kernels/moe/gather_gemv_iq2xxs.rs similarity index 96% rename from crates/metaltile-std/src/ffai/moe_gather_gemv_iq2xxs.rs rename to crates/metaltile-std/src/kernels/moe/gather_gemv_iq2xxs.rs index 8b37ec89..48c228e8 100644 --- a/crates/metaltile-std/src/ffai/moe_gather_gemv_iq2xxs.rs +++ b/crates/metaltile-std/src/kernels/moe/gather_gemv_iq2xxs.rs @@ -48,7 +48,7 @@ use metaltile::kernel; #[kernel] -pub fn ffai_moe_gather_gemv_iq2xxs( +pub fn mt_moe_gather_gemv_iq2xxs( x: Tensor, qs_all: Tensor, d_all: Tensor, @@ -113,12 +113,12 @@ pub fn ffai_moe_gather_gemv_iq2xxs( pub mod kernel_tests { use metaltile::test::*; - use super::ffai_moe_gather_gemv_iq2xxs; + use super::mt_moe_gather_gemv_iq2xxs; #[test] fn codegen_gather_gemv_iq2xxs_smoke() { for dt in [DType::F32, DType::F16, DType::BF16] { - let ir = ffai_moe_gather_gemv_iq2xxs::kernel_ir_for(dt); + let ir = mt_moe_gather_gemv_iq2xxs::kernel_ir_for(dt); assert!(!ir.body.ops.is_empty(), "no ops for {dt:?}"); assert!(ir.params.iter().any(|p| p.name == "qs_all"), "missing qs_all"); assert!(ir.params.iter().any(|p| p.name == "grid"), "missing grid"); @@ -129,7 +129,7 @@ pub mod kernel_tests { pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_gather_gemv_iq2xxs; + use super::mt_moe_gather_gemv_iq2xxs; // n_slots=6 routed experts; production gate/up dims (m_out=2048, k_in=4096). #[bench(dtypes = [f32, f16, bf16])] @@ -138,7 +138,7 @@ pub mod kernel_benches { let m_out = 2048usize; let k_in = 4096usize; let nblk = m_out * (k_in / 256); - BenchSetup::new(ffai_moe_gather_gemv_iq2xxs::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gather_gemv_iq2xxs::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", k_in, dt)) .buffer(BenchBuffer::random("qs_all", n_slots * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/ffai/moe_gemv_rows_iq2xxs.rs b/crates/metaltile-std/src/kernels/moe/gemv_rows_iq2xxs.rs similarity index 93% rename from crates/metaltile-std/src/ffai/moe_gemv_rows_iq2xxs.rs rename to crates/metaltile-std/src/kernels/moe/gemv_rows_iq2xxs.rs index 9facade7..484f2a5e 100644 --- a/crates/metaltile-std/src/ffai/moe_gemv_rows_iq2xxs.rs +++ b/crates/metaltile-std/src/kernels/moe/gemv_rows_iq2xxs.rs @@ -1,9 +1,9 @@ //! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric //! SPDX-License-Identifier: Apache-2.0 //! Prefill MoE IQ2_XXS GEMV-over-rows — the fast decode gemv -//! (ffai_moe_gather_gemv_iq2xxs, ~270 GB/s) applied to a whole batch of +//! (mt_moe_gather_gemv_iq2xxs, ~270 GB/s) applied to a whole batch of //! M = N*topK (token,expert) rows in ONE dispatch, instead of the -//! coop-tile MMA bgemm (ffai_moe_bgemm_iq2xxs_mpp, ~4-10 GB/s — 15-70x +//! coop-tile MMA bgemm (mt_moe_bgemm_iq2xxs_mpp, ~4-10 GB/s — 15-70x //! slower). The bgemm's MMA staging + barriers dominate at these quant //! shapes; the direct simd_sum dot-product the gemv uses is far faster. //! @@ -22,7 +22,7 @@ use metaltile::kernel; #[kernel] -pub fn ffai_moe_gemv_rows_iq2xxs( +pub fn mt_moe_gemv_rows_iq2xxs( x: Tensor, qs_all: Tensor, d_all: Tensor, @@ -81,7 +81,7 @@ pub fn ffai_moe_gemv_rows_iq2xxs( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_gemv_rows_iq2xxs; + use super::mt_moe_gemv_rows_iq2xxs; // M=256 rows, production gate/up dims (m_out=2048, k_in=4096). #[bench(dtypes = [f32, f16, bf16])] @@ -91,7 +91,7 @@ pub mod kernel_benches { let m_out = 2048usize; let k_in = 4096usize; let nblk = m_out * (k_in / 256); - BenchSetup::new(ffai_moe_gemv_rows_iq2xxs::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gemv_rows_iq2xxs::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", m_total * k_in, dt)) .buffer(BenchBuffer::random("qs_all", n_experts * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/ffai/moe_gemv_rows_q2k.rs b/crates/metaltile-std/src/kernels/moe/gemv_rows_q2k.rs similarity index 95% rename from crates/metaltile-std/src/ffai/moe_gemv_rows_q2k.rs rename to crates/metaltile-std/src/kernels/moe/gemv_rows_q2k.rs index 6d3db4c2..4c507d37 100644 --- a/crates/metaltile-std/src/ffai/moe_gemv_rows_q2k.rs +++ b/crates/metaltile-std/src/kernels/moe/gemv_rows_q2k.rs @@ -1,7 +1,7 @@ //! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric //! SPDX-License-Identifier: Apache-2.0 //! Prefill MoE Q2_K GEMV-over-rows (down projection) — the Q2_K twin of -//! ffai_moe_gemv_rows_iq2xxs. Replaces the slow coop-tile bgemm +//! mt_moe_gemv_rows_iq2xxs. Replaces the slow coop-tile bgemm //! (gather_bgemm_q2k_mpp ~10 GB/s) with the fast decode-style direct //! simd_sum dot-product applied to a whole batch of M=(token,expert) rows //! in ONE dispatch. Reads the resident split pool (qs u32 / scales u8 / @@ -15,7 +15,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_gemv_rows_q2k( +pub fn mt_moe_gemv_rows_q2k( x: Tensor, qs: Tensor, scales: Tensor, @@ -80,7 +80,7 @@ pub fn ffai_moe_gemv_rows_q2k( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_gemv_rows_q2k; + use super::mt_moe_gemv_rows_q2k; // M=256 rows, production down dims (m_out=4096 hidden, k_in=2048 intermediate). #[bench(dtypes = [f32, f16, bf16])] @@ -90,7 +90,7 @@ pub mod kernel_benches { let m_out = 4096usize; let k_in = 2048usize; let nblk = m_out * (k_in / 256); - BenchSetup::new(ffai_moe_gemv_rows_q2k::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gemv_rows_q2k::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", m_total * k_in, dt)) .buffer(BenchBuffer::random("qs", n_experts * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/ffai/moe_gemv_rows_view_iq2xxs.rs b/crates/metaltile-std/src/kernels/moe/gemv_rows_view_iq2xxs.rs similarity index 96% rename from crates/metaltile-std/src/ffai/moe_gemv_rows_view_iq2xxs.rs rename to crates/metaltile-std/src/kernels/moe/gemv_rows_view_iq2xxs.rs index d18088dd..c88f8120 100644 --- a/crates/metaltile-std/src/ffai/moe_gemv_rows_view_iq2xxs.rs +++ b/crates/metaltile-std/src/kernels/moe/gemv_rows_view_iq2xxs.rs @@ -16,7 +16,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_gemv_rows_view_iq2xxs( +pub fn mt_moe_gemv_rows_view_iq2xxs( x: Tensor, view_u8: Tensor, grid: Tensor, @@ -98,7 +98,7 @@ pub fn ffai_moe_gemv_rows_view_iq2xxs( /// gemv's ~100 GB/s (the u8 path was ~3.5). If so, pool-elimination is viable. #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_gemv_rows_view_u16_iq2xxs( +pub fn mt_moe_gemv_rows_view_u16_iq2xxs( x: Tensor, view_u16: Tensor, grid: Tensor, @@ -167,7 +167,7 @@ pub fn ffai_moe_gemv_rows_view_u16_iq2xxs( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_gemv_rows_view_u16_iq2xxs; + use super::mt_moe_gemv_rows_view_u16_iq2xxs; #[bench(dtypes = [f32, f16, bf16])] fn bench_gemv_rows_view_u16_iq2xxs(dt: DType) -> BenchSetup { @@ -177,7 +177,7 @@ pub mod kernel_benches { let k_in = 4096usize; let nblk = m_out * (k_in / 256); let view_bytes = n_experts * nblk * 66; - BenchSetup::new(ffai_moe_gemv_rows_view_u16_iq2xxs::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gemv_rows_view_u16_iq2xxs::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", m_total * k_in, dt)) .buffer(BenchBuffer::random("view_u16", view_bytes / 2, DType::U16)) @@ -198,7 +198,7 @@ pub mod kernel_benches { pub mod kernel_benches_old { use metaltile::{bench, test::*}; - use super::ffai_moe_gemv_rows_view_iq2xxs; + use super::mt_moe_gemv_rows_view_iq2xxs; #[bench(dtypes = [f32, f16, bf16])] fn bench_gemv_rows_view_iq2xxs(dt: DType) -> BenchSetup { @@ -208,7 +208,7 @@ pub mod kernel_benches_old { let k_in = 4096usize; let nblk = m_out * (k_in / 256); let view_bytes = n_experts * nblk * 66; - BenchSetup::new(ffai_moe_gemv_rows_view_iq2xxs::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gemv_rows_view_iq2xxs::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", m_total * k_in, dt)) .buffer(BenchBuffer::random("view_u8", view_bytes, DType::U8)) diff --git a/crates/metaltile-std/src/ffai/moe_gemv_ws_iq2xxs.rs b/crates/metaltile-std/src/kernels/moe/gemv_ws_iq2xxs.rs similarity index 97% rename from crates/metaltile-std/src/ffai/moe_gemv_ws_iq2xxs.rs rename to crates/metaltile-std/src/kernels/moe/gemv_ws_iq2xxs.rs index 0c92b40f..e96fec09 100644 --- a/crates/metaltile-std/src/ffai/moe_gemv_ws_iq2xxs.rs +++ b/crates/metaltile-std/src/kernels/moe/gemv_ws_iq2xxs.rs @@ -27,7 +27,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_gemv_ws_iq2xxs( +pub fn mt_moe_gemv_ws_iq2xxs( x: Tensor, qs_all: Tensor, d_all: Tensor, @@ -123,7 +123,7 @@ pub fn ffai_moe_gemv_ws_iq2xxs( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_gemv_ws_iq2xxs; + use super::mt_moe_gemv_ws_iq2xxs; // M=256 rows, production gate/up dims, 8 rows/tile. #[bench(dtypes = [f32, f16, bf16])] @@ -134,7 +134,7 @@ pub mod kernel_benches { let k_in = 4096usize; let rows_per_tile = 8usize; let nblk = m_out * (k_in / 256); - BenchSetup::new(ffai_moe_gemv_ws_iq2xxs::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gemv_ws_iq2xxs::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", m_total * k_in, dt)) .buffer(BenchBuffer::random("qs_all", n_experts * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/ffai/moe_gemv_ws_q2k.rs b/crates/metaltile-std/src/kernels/moe/gemv_ws_q2k.rs similarity index 96% rename from crates/metaltile-std/src/ffai/moe_gemv_ws_q2k.rs rename to crates/metaltile-std/src/kernels/moe/gemv_ws_q2k.rs index 0b296667..58ac2b21 100644 --- a/crates/metaltile-std/src/ffai/moe_gemv_ws_q2k.rs +++ b/crates/metaltile-std/src/kernels/moe/gemv_ws_q2k.rs @@ -1,7 +1,7 @@ //! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric //! SPDX-License-Identifier: Apache-2.0 //! Prefill MoE Q2_K WEIGHT-STATIONARY gemv (down projection) — the Q2_K -//! twin of ffai_moe_gemv_ws_iq2xxs. Dequants each expert's weight row +//! twin of mt_moe_gemv_ws_iq2xxs. Dequants each expert's weight row //! W_down[expert,m,:] ONCE into threadgroup memory and reuses it across //! all rows of that expert in the tile (amortized like bm64 but at gemv //! speed). Rows are pre-permuted by expert (contiguous), so a tile is @@ -16,7 +16,7 @@ use metaltile::kernel; #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_moe_gemv_ws_q2k( +pub fn mt_moe_gemv_ws_q2k( x: Tensor, qs: Tensor, scales: Tensor, @@ -104,7 +104,7 @@ pub fn ffai_moe_gemv_ws_q2k( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_gemv_ws_q2k; + use super::mt_moe_gemv_ws_q2k; // M=256 rows, production down dims (m_out=4096 hidden, k_in=2048). #[bench(dtypes = [f32, f16, bf16])] @@ -115,7 +115,7 @@ pub mod kernel_benches { let k_in = 2048usize; let rows_per_tile = 8usize; let nblk = m_out * (k_in / 256); - BenchSetup::new(ffai_moe_gemv_ws_q2k::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_gemv_ws_q2k::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("x", m_total * k_in, dt)) .buffer(BenchBuffer::random("qs", n_experts * nblk * 16, DType::U32)) diff --git a/crates/metaltile-std/src/kernels/moe/mod.rs b/crates/metaltile-std/src/kernels/moe/mod.rs index 9e90d8e7..451c8d43 100644 --- a/crates/metaltile-std/src/kernels/moe/mod.rs +++ b/crates/metaltile-std/src/kernels/moe/mod.rs @@ -1,10 +1,63 @@ //! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric //! SPDX-License-Identifier: Apache-2.0 //! Mixture-of-experts kernels — the moe family (see -//! `docs/specs/KERNEL_CONSOLIDATION_PLAN.md`). Migrated ahead of the full moe -//! wave: the batched Q4 expert-gather projections (up / down / weighted-sum) -//! and the on-device router pre-score. The remaining moe_* kernels still live -//! in `ffai/` and land here when the moe family is consolidated. +//! `docs/specs/KERNEL_CONSOLIDATION_PLAN.md`). The routing layer (top-k expert +//! selection + permute/unpermute + router pre-scores), the per-expert quantized +//! matmul cells in their various forms (MPP grouped BGEMM, GGUF q2k/iq2xxs +//! BGEMM/GEMV, batched Q4 gather), and the down-projection combine. Migrated +//! from the legacy `mlx/` + `ffai/` split. +//! +//! Filenames drop the redundant `moe_` prefix (the folder provides it); kernel +//! names keep `mt_moe_*`. The per-format `*_block_scaled` matrices move as-is; +//! the format-axis fold (plan §7) is deferred. `orchestration.rs` is large and +//! is slated for a follow-up split (router_topk / permute / gather_qmm). -pub mod gather_q4; +// Routing — top-k expert selection, permute/unpermute, router pre-scores. +pub mod orchestration; +pub mod router_topk_biased; +pub mod router_sigmoid_bias; +pub mod router_sqrtsoftplus; pub mod sigmoid_bias; + +// MPP grouped BGEMM (one ABI, tile-geometry / bit-width variants; shared +// test/bench helpers in `mpp_shared`). +pub mod mpp; +pub mod mpp_shared; +pub mod mpp_int8; +pub mod mpp_bm8; +pub mod mpp_bm8_int8; +pub mod mpp_bm64; +pub mod mpp_bm64_int8; +pub mod mpp_block_scaled; +pub mod mpp_bm8_block_scaled; +pub mod mpp_bm64_block_scaled; + +// GGUF-format per-expert matmul / matvec (q2k, iq2xxs, q4). +pub mod bgemm_q2k_bm64; +pub mod bgemm_q2k_mpp; +pub mod bgemm_q2k_view; +pub mod bgemm_q2k_view_u16_bm64; +pub mod bgemm_q4_bm64; +pub mod bgemm_iq2xxs_bm64; +pub mod bgemm_iq2xxs_mpp; +pub mod bgemm_iq2xxs_view; +pub mod bgemm_iq2xxs_view_u16_bm64; +pub mod gather_down_q2k; +pub mod gather_gemv_iq2xxs; +pub mod gemv_rows_q2k; +pub mod gemv_rows_iq2xxs; +pub mod gemv_rows_view_iq2xxs; +pub mod gemv_ws_q2k; +pub mod gemv_ws_iq2xxs; + +// Batched Q4 expert gather (up / down / weighted-sum), seeded from gemv_q8. +pub mod gather_q4; + +// Down-projection combine (swiglu-fused accumulate, weighted sum). +pub mod down_swiglu_accum; +pub mod down_weighted_sum_f16; + +// Expert-indexed dequant GEMV + block-scaled MoE matmul. +pub mod dequant_gemv_expert_indexed; +pub mod dequant_gemv_expert_indexed_block_scaled; +pub mod block_scaled; diff --git a/crates/metaltile-std/src/ffai/moe_mpp.rs b/crates/metaltile-std/src/kernels/moe/mpp.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_mpp.rs rename to crates/metaltile-std/src/kernels/moe/mpp.rs index 063c0d0f..43f0741d 100644 --- a/crates/metaltile-std/src/ffai/moe_mpp.rs +++ b/crates/metaltile-std/src/kernels/moe/mpp.rs @@ -250,7 +250,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int4_bm16_mpp; - use crate::ffai::moe_mpp_shared::{MmaTestShape, int4_indexed_setup}; + use crate::kernels::moe::mpp_shared::{MmaTestShape, int4_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int4_bm16_mpp(dt: DType) -> TestSetup { @@ -271,7 +271,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int4_bm16_mpp; - use crate::ffai::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int4_bm16_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/ffai/moe_mpp_block_scaled.rs b/crates/metaltile-std/src/kernels/moe/mpp_block_scaled.rs similarity index 100% rename from crates/metaltile-std/src/ffai/moe_mpp_block_scaled.rs rename to crates/metaltile-std/src/kernels/moe/mpp_block_scaled.rs diff --git a/crates/metaltile-std/src/ffai/moe_mpp_bm64.rs b/crates/metaltile-std/src/kernels/moe/mpp_bm64.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_mpp_bm64.rs rename to crates/metaltile-std/src/kernels/moe/mpp_bm64.rs index a0269338..f542d03a 100644 --- a/crates/metaltile-std/src/ffai/moe_mpp_bm64.rs +++ b/crates/metaltile-std/src/kernels/moe/mpp_bm64.rs @@ -220,7 +220,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int4_bm64_mpp; - use crate::ffai::moe_mpp_shared::{MmaTestShape, int4_indexed_setup}; + use crate::kernels::moe::mpp_shared::{MmaTestShape, int4_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int4_bm64_mpp(dt: DType) -> TestSetup { @@ -241,7 +241,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int4_bm64_mpp; - use crate::ffai::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int4_bm64_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/ffai/moe_mpp_bm64_block_scaled.rs b/crates/metaltile-std/src/kernels/moe/mpp_bm64_block_scaled.rs similarity index 100% rename from crates/metaltile-std/src/ffai/moe_mpp_bm64_block_scaled.rs rename to crates/metaltile-std/src/kernels/moe/mpp_bm64_block_scaled.rs diff --git a/crates/metaltile-std/src/ffai/moe_mpp_bm64_int8.rs b/crates/metaltile-std/src/kernels/moe/mpp_bm64_int8.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_mpp_bm64_int8.rs rename to crates/metaltile-std/src/kernels/moe/mpp_bm64_int8.rs index f767dba7..b7605362 100644 --- a/crates/metaltile-std/src/ffai/moe_mpp_bm64_int8.rs +++ b/crates/metaltile-std/src/kernels/moe/mpp_bm64_int8.rs @@ -245,7 +245,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int8_bm64_mpp; - use crate::ffai::moe_mpp_shared::{MmaTestShape, int8_indexed_setup}; + use crate::kernels::moe::mpp_shared::{MmaTestShape, int8_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int8_bm64_mpp(dt: DType) -> TestSetup { @@ -270,7 +270,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int8_bm64_mpp; - use crate::ffai::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int8_bm64_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/ffai/moe_mpp_bm8.rs b/crates/metaltile-std/src/kernels/moe/mpp_bm8.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_mpp_bm8.rs rename to crates/metaltile-std/src/kernels/moe/mpp_bm8.rs index 0d40992a..80fc9e6e 100644 --- a/crates/metaltile-std/src/ffai/moe_mpp_bm8.rs +++ b/crates/metaltile-std/src/kernels/moe/mpp_bm8.rs @@ -214,7 +214,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int4_bm8_mpp; - use crate::ffai::moe_mpp_shared::{MmaTestShape, int4_indexed_setup}; + use crate::kernels::moe::mpp_shared::{MmaTestShape, int4_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int4_bm8_mpp(dt: DType) -> TestSetup { @@ -235,7 +235,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int4_bm8_mpp; - use crate::ffai::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int4_bm8_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/ffai/moe_mpp_bm8_block_scaled.rs b/crates/metaltile-std/src/kernels/moe/mpp_bm8_block_scaled.rs similarity index 100% rename from crates/metaltile-std/src/ffai/moe_mpp_bm8_block_scaled.rs rename to crates/metaltile-std/src/kernels/moe/mpp_bm8_block_scaled.rs diff --git a/crates/metaltile-std/src/ffai/moe_mpp_bm8_int8.rs b/crates/metaltile-std/src/kernels/moe/mpp_bm8_int8.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_mpp_bm8_int8.rs rename to crates/metaltile-std/src/kernels/moe/mpp_bm8_int8.rs index 3efd35cc..05ea0351 100644 --- a/crates/metaltile-std/src/ffai/moe_mpp_bm8_int8.rs +++ b/crates/metaltile-std/src/kernels/moe/mpp_bm8_int8.rs @@ -241,7 +241,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int8_bm8_mpp; - use crate::ffai::moe_mpp_shared::{MmaTestShape, int8_indexed_setup}; + use crate::kernels::moe::mpp_shared::{MmaTestShape, int8_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int8_bm8_mpp(dt: DType) -> TestSetup { @@ -265,7 +265,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int8_bm8_mpp; - use crate::ffai::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int8_bm8_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/ffai/moe_mpp_int8.rs b/crates/metaltile-std/src/kernels/moe/mpp_int8.rs similarity index 98% rename from crates/metaltile-std/src/ffai/moe_mpp_int8.rs rename to crates/metaltile-std/src/kernels/moe/mpp_int8.rs index e18f3658..658afc8d 100644 --- a/crates/metaltile-std/src/ffai/moe_mpp_int8.rs +++ b/crates/metaltile-std/src/kernels/moe/mpp_int8.rs @@ -257,7 +257,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int8_bm16_mpp; - use crate::ffai::moe_mpp_shared::{MmaTestShape, int8_indexed_setup}; + use crate::kernels::moe::mpp_shared::{MmaTestShape, int8_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int8_bm16_mpp(dt: DType) -> TestSetup { @@ -281,7 +281,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int8_bm16_mpp; - use crate::ffai::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int8_bm16_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/ffai/moe_mpp_shared.rs b/crates/metaltile-std/src/kernels/moe/mpp_shared.rs similarity index 100% rename from crates/metaltile-std/src/ffai/moe_mpp_shared.rs rename to crates/metaltile-std/src/kernels/moe/mpp_shared.rs diff --git a/crates/metaltile-std/src/ffai/moe.rs b/crates/metaltile-std/src/kernels/moe/orchestration.rs similarity index 100% rename from crates/metaltile-std/src/ffai/moe.rs rename to crates/metaltile-std/src/kernels/moe/orchestration.rs diff --git a/crates/metaltile-std/src/ffai/moe_router_sigmoid_bias.rs b/crates/metaltile-std/src/kernels/moe/router_sigmoid_bias.rs similarity index 94% rename from crates/metaltile-std/src/ffai/moe_router_sigmoid_bias.rs rename to crates/metaltile-std/src/kernels/moe/router_sigmoid_bias.rs index dd9f46e2..a35f901d 100644 --- a/crates/metaltile-std/src/ffai/moe_router_sigmoid_bias.rs +++ b/crates/metaltile-std/src/kernels/moe/router_sigmoid_bias.rs @@ -44,7 +44,7 @@ use metaltile::kernel; // a no-`` signature. The new declarative `#[bench]` on // `kernel_benches::bench_router` below handles registration directly. #[kernel] -pub fn ffai_moe_router_sigmoid_bias( +pub fn mt_moe_router_sigmoid_bias( logits: Tensor, bias: Tensor, mut scores: Tensor, @@ -62,7 +62,7 @@ pub fn ffai_moe_router_sigmoid_bias( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::ffai_moe_router_sigmoid_bias; + use super::mt_moe_router_sigmoid_bias; use crate::utils::{pack_f32, unpack_f32}; fn setup(n_experts: usize) -> TestSetup { @@ -73,7 +73,7 @@ pub mod kernel_tests { let b_dt = unpack_f32(&pack_f32(&bias, dt), dt); let expected: Vec = l_dt.iter().zip(&b_dt).map(|(&l, &b)| 1.0_f32 / (1.0 + (-l).exp()) + b).collect(); - TestSetup::new(ffai_moe_router_sigmoid_bias::kernel_ir()) + TestSetup::new(mt_moe_router_sigmoid_bias::kernel_ir()) .input(TestBuffer::from_vec("logits", pack_f32(&logits, dt), dt)) .input(TestBuffer::from_vec("bias", pack_f32(&bias, dt), dt)) .input(TestBuffer::zeros("scores", n_experts, dt)) @@ -96,13 +96,13 @@ pub mod kernel_tests { pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_router_sigmoid_bias; + use super::mt_moe_router_sigmoid_bias; #[bench(dtypes = [f32])] fn bench_router(_dt: DType) -> BenchSetup { let dt = DType::F32; let n_experts = 288usize; - BenchSetup::new(ffai_moe_router_sigmoid_bias::kernel_ir()) + BenchSetup::new(mt_moe_router_sigmoid_bias::kernel_ir()) .buffer(BenchBuffer::random("logits", n_experts, dt)) .buffer(BenchBuffer::random("bias", n_experts, dt)) .buffer(BenchBuffer::zeros("scores", n_experts, dt).output()) diff --git a/crates/metaltile-std/src/ffai/moe_router_sqrtsoftplus.rs b/crates/metaltile-std/src/kernels/moe/router_sqrtsoftplus.rs similarity index 95% rename from crates/metaltile-std/src/ffai/moe_router_sqrtsoftplus.rs rename to crates/metaltile-std/src/kernels/moe/router_sqrtsoftplus.rs index d249cd7b..d176a738 100644 --- a/crates/metaltile-std/src/ffai/moe_router_sqrtsoftplus.rs +++ b/crates/metaltile-std/src/kernels/moe/router_sqrtsoftplus.rs @@ -17,7 +17,7 @@ //! downstream top-k + normalize + scale chain consumes whichever it //! needs without re-running the scoring math. //! -//! Compared to the sister `ffai_moe_router_sigmoid_bias`: +//! Compared to the sister `mt_moe_router_sigmoid_bias`: //! - sigmoid+bias: `s = sigmoid(x)`; bounded in (0, 1). //! - sqrtsoftplus: `s = sqrt(log(1 + exp(x)))`; unbounded, larger //! dynamic range — paired with the bias-correction for selection. @@ -45,7 +45,7 @@ use metaltile::kernel; // legacy `bench(...)` shape; declarative `#[bench]` below registers // for `tile bench`. #[kernel] -pub fn ffai_moe_router_sqrtsoftplus( +pub fn mt_moe_router_sqrtsoftplus( logits: Tensor, bias: Tensor, mut score_unbiased: Tensor, @@ -82,7 +82,7 @@ pub fn ffai_moe_router_sqrtsoftplus( pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::ffai_moe_router_sqrtsoftplus; + use super::mt_moe_router_sqrtsoftplus; use crate::utils::pack_f32; /// CPU reference. Mirrors the GPU kernel for tight-tolerance check. @@ -113,7 +113,7 @@ pub mod kernel_tests { fn setup(n_experts: usize) -> TestSetup { let dt = DType::F32; let (logits, bias, score_unbiased, score_biased) = cpu_reference(n_experts); - TestSetup::new(ffai_moe_router_sqrtsoftplus::kernel_ir()) + TestSetup::new(mt_moe_router_sqrtsoftplus::kernel_ir()) .input(TestBuffer::from_vec("logits", pack_f32(&logits, dt), dt)) .input(TestBuffer::from_vec("bias", pack_f32(&bias, dt), dt)) .input(TestBuffer::zeros("score_unbiased", n_experts, dt)) @@ -139,13 +139,13 @@ pub mod kernel_tests { pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::ffai_moe_router_sqrtsoftplus; + use super::mt_moe_router_sqrtsoftplus; #[bench(dtypes = [f32])] fn bench_router(_dt: DType) -> BenchSetup { let dt = DType::F32; let n_experts = 288usize; - BenchSetup::new(ffai_moe_router_sqrtsoftplus::kernel_ir()) + BenchSetup::new(mt_moe_router_sqrtsoftplus::kernel_ir()) .buffer(BenchBuffer::random("logits", n_experts, dt)) .buffer(BenchBuffer::random("bias", n_experts, dt)) .buffer(BenchBuffer::zeros("score_unbiased", n_experts, dt).output()) diff --git a/crates/metaltile-std/src/ffai/dsv4_router_topk.rs b/crates/metaltile-std/src/kernels/moe/router_topk_biased.rs similarity index 95% rename from crates/metaltile-std/src/ffai/dsv4_router_topk.rs rename to crates/metaltile-std/src/kernels/moe/router_topk_biased.rs index 418edd8c..0b21c5b5 100644 --- a/crates/metaltile-std/src/ffai/dsv4_router_topk.rs +++ b/crates/metaltile-std/src/kernels/moe/router_topk_biased.rs @@ -26,7 +26,7 @@ use metaltile::kernel; #[kernel] -pub fn mt_dsv4_router_topk( +pub fn mt_moe_router_topk_biased( score_biased: Tensor, score_unbiased: Tensor, mut indices_out: Tensor, @@ -104,13 +104,13 @@ pub fn mt_remap_u32( pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::{mt_dsv4_router_topk, mt_remap_u32}; + use super::{mt_moe_router_topk_biased, mt_remap_u32}; #[bench(dtypes = [f32, f16, bf16])] - fn bench_dsv4_router_topk(dt: DType) -> BenchSetup { + fn bench_moe_router_topk_biased(dt: DType) -> BenchSetup { let n_experts = 256usize; let k = 6usize; - BenchSetup::new(mt_dsv4_router_topk::kernel_ir_for(dt)) + BenchSetup::new(mt_moe_router_topk_biased::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("score_biased", n_experts, dt)) .buffer(BenchBuffer::random("score_unbiased", n_experts, dt)) diff --git a/crates/metaltile-std/src/kernels/moe/sigmoid_bias.rs b/crates/metaltile-std/src/kernels/moe/sigmoid_bias.rs index 9a2ccb33..416e1677 100644 --- a/crates/metaltile-std/src/kernels/moe/sigmoid_bias.rs +++ b/crates/metaltile-std/src/kernels/moe/sigmoid_bias.rs @@ -7,7 +7,7 @@ use metaltile::kernel; /// MoE router pre-scores (NemotronH / DeepSeek-V3 noaux, sigmoid variant): /// `unbiased[i] = sigmoid(logit[i])`, `biased[i] = unbiased[i] + e_score_correction_bias[i]`. -/// Feeds `mt_dsv4_router_topk` (top-k by biased, weights from unbiased) so the whole +/// Feeds `mt_moe_router_topk_biased` (top-k by biased, weights from unbiased) so the whole /// router stays ON-DEVICE — no per-MoE-layer dl(gate)+host-topk+up(idx) sync round-trip. #[kernel] pub fn mt_moe_sigmoid_bias( diff --git a/crates/metaltile-std/src/mlx/mod.rs b/crates/metaltile-std/src/mlx/mod.rs index d98cd504..8047f0ee 100644 --- a/crates/metaltile-std/src/mlx/mod.rs +++ b/crates/metaltile-std/src/mlx/mod.rs @@ -14,10 +14,9 @@ //! expected to wire up eventually, it lives in `ffai/` until the //! comparison lands. -// block_scaled_dequant → quant family; block_scaled_moe → moe family -// (migrated later). The quantized matmuls moved to kernels/gemm/. +// block_scaled_dequant → quant family (migrated later); block_scaled_moe → +// kernels/moe/. The quantized matmuls moved to kernels/gemm/. pub mod block_scaled_dequant; -pub mod block_scaled_moe; pub mod scaled_dot_product_attention; pub mod sdpa_vector; pub mod steel; diff --git a/crates/metaltile-std/tests/bm64_vs_gemvrows_iq2xxs.rs b/crates/metaltile-std/tests/bm64_vs_gemvrows_iq2xxs.rs index ea5db4c5..6e9faac5 100644 --- a/crates/metaltile-std/tests/bm64_vs_gemvrows_iq2xxs.rs +++ b/crates/metaltile-std/tests/bm64_vs_gemvrows_iq2xxs.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! Direct comparison: ffai_moe_bgemm_iq2xxs_bm64 vs ffai_moe_gemv_rows_iq2xxs +//! Direct comparison: mt_moe_bgemm_iq2xxs_bm64 vs mt_moe_gemv_rows_iq2xxs //! on IDENTICAL pool/x/indices. Both claim to compute gateP[row,m] = //! W[expert(row),m,:]·x[row,:]; the prefill shows them disagreeing. This //! reproduces it in isolation. cosine should be ~1.0. @@ -12,9 +12,9 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{ - moe_bgemm_iq2xxs_bm64::ffai_moe_bgemm_iq2xxs_bm64, - moe_gemv_rows_iq2xxs::ffai_moe_gemv_rows_iq2xxs, +use metaltile_std::kernels::moe::{ + bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, + gemv_rows_iq2xxs::mt_moe_gemv_rows_iq2xxs, }; fn xorshift(s: &mut u32) -> u32 { @@ -67,7 +67,7 @@ fn bm64_matches_gemvrows() { b.insert("m_total".into(), (t_rows as u32).to_le_bytes().to_vec()); b.insert("n_out".into(), (n_out as u32).to_le_bytes().to_vec()); b.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); - let mut kb = ffai_moe_bgemm_iq2xxs_bm64::kernel_ir_for(Dt::F32.to_dtype()); + let mut kb = mt_moe_bgemm_iq2xxs_bm64::kernel_ir_for(Dt::F32.to_dtype()); kb.mode = KernelMode::Reduction; let rb = ctx .dispatch_with_grid(&kb, &b, &BTreeMap::new(), [n_out / 64, t_rows.div_ceil(64), 1], [ @@ -88,7 +88,7 @@ fn bm64_matches_gemvrows() { g.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); g.insert("m_out".into(), (n_out as u32).to_le_bytes().to_vec()); g.insert("m_total".into(), (t_rows as u32).to_le_bytes().to_vec()); - let mut kg = ffai_moe_gemv_rows_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); + let mut kg = mt_moe_gemv_rows_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); kg.mode = KernelMode::Reduction; let rg = ctx.dispatch_with_grid(&kg, &g, &BTreeMap::new(), [n_out, t_rows, 1], [32, 1, 1]).unwrap(); diff --git a/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs b/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs index ae82222b..c7b80f0a 100644 --- a/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs +++ b/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::mt_dsv4_router_topk` — top-K by biased +//! GPU correctness for `ffai::mt_moe_router_topk_biased` — top-K by biased //! score, weights = unbiased[chosen] renormalised to sum 1. #![cfg(target_os = "macos")] @@ -10,7 +10,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes, unpack_u32_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::dsv4_router_topk::{mt_dsv4_router_topk, mt_remap_u32}; +use metaltile_std::kernels::moe::router_topk_biased::{mt_moe_router_topk_biased, mt_remap_u32}; #[test] fn dsv4_router_topk_f32() { @@ -37,7 +37,7 @@ fn dsv4_router_topk_f32() { buffers.insert("k".into(), (k as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("ctx"); - let mut kernel = mt_dsv4_router_topk::kernel_ir_for(Dt::F32.to_dtype()); + let mut kernel = mt_moe_router_topk_biased::kernel_ir_for(Dt::F32.to_dtype()); kernel.mode = KernelMode::Reduction; let result = ctx .dispatch_with_grid(&kernel, &buffers, &BTreeMap::new(), [1, 1, 1], [32, 1, 1]) diff --git a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_bm64_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_bm64_correctness.rs index cd8a9867..8eb6bb66 100644 --- a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_bm64_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_bm64_correctness.rs @@ -1,7 +1,7 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 //! GPU correctness for the high-throughput bm64 IQ2_XXS BGEMM — must match -//! the proven 16×32 pool kernel (ffai_moe_gather_bgemm_iq2xxs_mpp) on +//! the proven 16×32 pool kernel (mt_moe_gather_bgemm_iq2xxs_mpp) on //! identical weights/x/indices (same dequant, only the tile geometry differs). //! cosine ≥ 0.999. NO 86GB model load. #![cfg(target_os = "macos")] @@ -12,9 +12,9 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{ - moe_bgemm_iq2xxs_bm64::ffai_moe_bgemm_iq2xxs_bm64, - moe_bgemm_iq2xxs_mpp::ffai_moe_gather_bgemm_iq2xxs_mpp, +use metaltile_std::kernels::moe::{ + bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, + bgemm_iq2xxs_mpp::mt_moe_gather_bgemm_iq2xxs_mpp, }; fn xorshift(s: &mut u32) -> u32 { @@ -66,7 +66,7 @@ fn bgemm_iq2xxs_bm64_matches_pool_kernel() { pool.insert("m_total".into(), (t_rows as u32).to_le_bytes().to_vec()); pool.insert("n_out".into(), (n_out as u32).to_le_bytes().to_vec()); pool.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); - let mut kp = ffai_moe_gather_bgemm_iq2xxs_mpp::kernel_ir_for(Dt::F32.to_dtype()); + let mut kp = mt_moe_gather_bgemm_iq2xxs_mpp::kernel_ir_for(Dt::F32.to_dtype()); kp.mode = KernelMode::Reduction; let rp = ctx .dispatch_with_grid(&kp, &pool, &BTreeMap::new(), [n_out / 32, t_rows.div_ceil(16), 1], [ @@ -78,7 +78,7 @@ fn bgemm_iq2xxs_bm64_matches_pool_kernel() { // bm64 kernel (same buffers; 128 threads, n_out/64 × M/64 grid). let mut b: BTreeMap> = pool.clone(); b.insert("out".into(), pack_bytes(&vec![0.0f32; t_rows * n_out], Dt::F32)); - let mut kb = ffai_moe_bgemm_iq2xxs_bm64::kernel_ir_for(Dt::F32.to_dtype()); + let mut kb = mt_moe_bgemm_iq2xxs_bm64::kernel_ir_for(Dt::F32.to_dtype()); kb.mode = KernelMode::Reduction; let rb = ctx .dispatch_with_grid(&kb, &b, &BTreeMap::new(), [n_out / 64, t_rows.div_ceil(64), 1], [ diff --git a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_mpp_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_mpp_correctness.rs index fd875e2a..775339e1 100644 --- a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_mpp_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_mpp_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_moe_gather_bgemm_iq2xxs_mpp` — the +//! GPU correctness for `ffai::mt_moe_gather_bgemm_iq2xxs_mpp` — the //! prefill IQ2_XXS grouped BGEMM. Oracle: per-row IQ2_XXS dequant gemv //! (same formula as ffai_gguf_dequant_iq2_xxs). Cosine ≥ 0.99 (MMA //! accumulation order differs from the scalar oracle). @@ -12,7 +12,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe_bgemm_iq2xxs_mpp::ffai_moe_gather_bgemm_iq2xxs_mpp; +use metaltile_std::kernels::moe::bgemm_iq2xxs_mpp::mt_moe_gather_bgemm_iq2xxs_mpp; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; @@ -97,7 +97,7 @@ fn bgemm_iq2xxs_mpp_matches_gemv_oracle() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = ffai_moe_gather_bgemm_iq2xxs_mpp::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = mt_moe_gather_bgemm_iq2xxs_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid(&k, &buffers, &BTreeMap::new(), [n_out / 32, t_rows.div_ceil(16), 1], [ diff --git a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_view_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_view_correctness.rs index b53e5a25..a1cd0a8e 100644 --- a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_view_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_view_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_moe_bgemm_iq2xxs_view` — the ZERO-COPY +//! GPU correctness for `ffai::mt_moe_bgemm_iq2xxs_view` — the ZERO-COPY //! prefill IQ2_XXS grouped BGEMM that reads raw 66-byte IQ2_XXS blocks //! straight from a no-copy mmap VIEW buffer (vs the repacked qs/d_f32 pool). //! Same oracle as the pool kernel (per-row IQ2_XXS dequant gemv), but the @@ -19,7 +19,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use half::f16; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe_bgemm_iq2xxs_view::ffai_moe_bgemm_iq2xxs_view; +use metaltile_std::kernels::moe::bgemm_iq2xxs_view::mt_moe_bgemm_iq2xxs_view; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; @@ -130,7 +130,7 @@ fn bgemm_iq2xxs_view_matches_gemv_oracle() { buffers.insert("expert_byte_stride".into(), (expert_byte_stride as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = ffai_moe_bgemm_iq2xxs_view::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = mt_moe_bgemm_iq2xxs_view::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid(&k, &buffers, &BTreeMap::new(), [n_out / 32, t_rows.div_ceil(16), 1], [ @@ -258,7 +258,7 @@ fn bgemm_iq2xxs_view_prod_dims_nonzero_offset() { buffers.insert("expert_byte_stride".into(), (expert_byte_stride as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = ffai_moe_bgemm_iq2xxs_view::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = mt_moe_bgemm_iq2xxs_view::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid(&k, &buffers, &BTreeMap::new(), [n_out / 32, t_rows.div_ceil(16), 1], [ diff --git a/crates/metaltile-std/tests/moe_bgemm_q2k_bm64_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_q2k_bm64_correctness.rs index 6e5073a9..7a7eef9d 100644 --- a/crates/metaltile-std/tests/moe_bgemm_q2k_bm64_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_q2k_bm64_correctness.rs @@ -1,7 +1,7 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 //! bm64 Q2_K BGEMM must match the proven 16×32 pool kernel -//! (ffai_moe_gather_bgemm_q2k_mpp). cosine ≥ 0.999. NO 86GB model load. +//! (mt_moe_gather_bgemm_q2k_mpp). cosine ≥ 0.999. NO 86GB model load. #![cfg(target_os = "macos")] mod common; @@ -10,9 +10,9 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{ - moe_bgemm_q2k_bm64::ffai_moe_bgemm_q2k_bm64, - moe_bgemm_q2k_mpp::ffai_moe_gather_bgemm_q2k_mpp, +use metaltile_std::kernels::moe::{ + bgemm_q2k_bm64::mt_moe_bgemm_q2k_bm64, + bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp, }; fn xorshift(s: &mut u32) -> u32 { @@ -64,7 +64,7 @@ fn bgemm_q2k_bm64_matches_pool_kernel() { pool.insert("m_total".into(), (t_rows as u32).to_le_bytes().to_vec()); pool.insert("n_out".into(), (n_out as u32).to_le_bytes().to_vec()); pool.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); - let mut kp = ffai_moe_gather_bgemm_q2k_mpp::kernel_ir_for(Dt::F32.to_dtype()); + let mut kp = mt_moe_gather_bgemm_q2k_mpp::kernel_ir_for(Dt::F32.to_dtype()); kp.mode = KernelMode::Reduction; let rp = ctx .dispatch_with_grid(&kp, &pool, &BTreeMap::new(), [n_out / 32, t_rows.div_ceil(16), 1], [ @@ -75,7 +75,7 @@ fn bgemm_q2k_bm64_matches_pool_kernel() { let mut b = pool.clone(); b.insert("out".into(), pack_bytes(&vec![0.0f32; t_rows * n_out], Dt::F32)); - let mut kb = ffai_moe_bgemm_q2k_bm64::kernel_ir_for(Dt::F32.to_dtype()); + let mut kb = mt_moe_bgemm_q2k_bm64::kernel_ir_for(Dt::F32.to_dtype()); kb.mode = KernelMode::Reduction; let rb = ctx .dispatch_with_grid(&kb, &b, &BTreeMap::new(), [n_out / 64, t_rows.div_ceil(64), 1], [ diff --git a/crates/metaltile-std/tests/moe_bgemm_q2k_mpp_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_q2k_mpp_correctness.rs index 9771fde7..5d7119b6 100644 --- a/crates/metaltile-std/tests/moe_bgemm_q2k_mpp_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_q2k_mpp_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_moe_gather_bgemm_q2k_mpp` — prefill Q2_K +//! GPU correctness for `ffai::mt_moe_gather_bgemm_q2k_mpp` — prefill Q2_K //! grouped BGEMM. Oracle: per-row Q2_K dequant gemv. Cosine ≥ 0.99. #![cfg(target_os = "macos")] @@ -10,7 +10,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe_bgemm_q2k_mpp::ffai_moe_gather_bgemm_q2k_mpp; +use metaltile_std::kernels::moe::bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp; // Shared Q2_K output-index → (qs byte, 2-bit shift) map (see PR #264/#265): the // kernel, quantizer, and this oracle all read the one definition in quant::gguf. use metaltile_std::quant::gguf::q2_k_qpos; @@ -96,7 +96,7 @@ fn bgemm_q2k_mpp_matches_gemv_oracle() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = ffai_moe_gather_bgemm_q2k_mpp::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = mt_moe_gather_bgemm_q2k_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid(&k, &buffers, &BTreeMap::new(), [n_out / 32, t_rows.div_ceil(16), 1], [ diff --git a/crates/metaltile-std/tests/moe_bgemm_q2k_view_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_q2k_view_correctness.rs index 81e1fb8a..46b0979d 100644 --- a/crates/metaltile-std/tests/moe_bgemm_q2k_view_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_q2k_view_correctness.rs @@ -1,10 +1,10 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_moe_bgemm_q2k_view` — the ZERO-COPY Q2_K +//! GPU correctness for `ffai::mt_moe_bgemm_q2k_view` — the ZERO-COPY Q2_K //! grouped BGEMM that reads raw 84-byte Q2_K blocks straight from a no-copy //! mmap VIEW buffer. Rather than re-derive a scalar oracle for the canonical //! Q2_K layout, this proves the view kernel produces the SAME output as the -//! PROVEN pool kernel (`ffai_moe_gather_bgemm_q2k_mpp`, the validated "Tokyo" +//! PROVEN pool kernel (`mt_moe_gather_bgemm_q2k_mpp`, the validated "Tokyo" //! path) on identical logical weights — the view just reads the raw bytes the //! pool would have been repacked from. cosine ≥ 0.999. NO 86 GB model load. #![cfg(target_os = "macos")] @@ -16,9 +16,9 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use half::f16; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{ - moe_bgemm_q2k_mpp::ffai_moe_gather_bgemm_q2k_mpp, - moe_bgemm_q2k_view::ffai_moe_bgemm_q2k_view, +use metaltile_std::kernels::moe::{ + bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp, + bgemm_q2k_view::mt_moe_bgemm_q2k_view, }; fn xorshift(s: &mut u32) -> u32 { @@ -79,7 +79,7 @@ fn bgemm_q2k_view_matches_pool_kernel() { pool.insert("n_out".into(), (n_out as u32).to_le_bytes().to_vec()); pool.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut kp = ffai_moe_gather_bgemm_q2k_mpp::kernel_ir_for(Dt::F32.to_dtype()); + let mut kp = mt_moe_gather_bgemm_q2k_mpp::kernel_ir_for(Dt::F32.to_dtype()); kp.mode = KernelMode::Reduction; let rp = ctx .dispatch_with_grid(&kp, &pool, &BTreeMap::new(), [n_out / 32, t_rows.div_ceil(16), 1], [ @@ -121,7 +121,7 @@ fn bgemm_q2k_view_matches_pool_kernel() { vb.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); vb.insert("tensor_byte_off".into(), 0u32.to_le_bytes().to_vec()); vb.insert("expert_byte_stride".into(), (expert_byte_stride as u32).to_le_bytes().to_vec()); - let mut kv = ffai_moe_bgemm_q2k_view::kernel_ir_for(Dt::F32.to_dtype()); + let mut kv = mt_moe_bgemm_q2k_view::kernel_ir_for(Dt::F32.to_dtype()); kv.mode = KernelMode::Reduction; let rv = ctx .dispatch_with_grid(&kv, &vb, &BTreeMap::new(), [n_out / 32, t_rows.div_ceil(16), 1], [ diff --git a/crates/metaltile-std/tests/moe_bm64_ragged_correctness.rs b/crates/metaltile-std/tests/moe_bm64_ragged_correctness.rs index b4d11f7f..b3c42ca3 100644 --- a/crates/metaltile-std/tests/moe_bm64_ragged_correctness.rs +++ b/crates/metaltile-std/tests/moe_bm64_ragged_correctness.rs @@ -21,9 +21,9 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{ - moe_bgemm_iq2xxs_bm64::ffai_moe_bgemm_iq2xxs_bm64, - moe_gemv_rows_iq2xxs::ffai_moe_gemv_rows_iq2xxs, +use metaltile_std::kernels::moe::{ + bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, + gemv_rows_iq2xxs::mt_moe_gemv_rows_iq2xxs, }; fn read_u32(p: &str) -> Vec { @@ -155,7 +155,7 @@ fn run_bm64( buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("ctx"); - let mut k = ffai_moe_bgemm_iq2xxs_bm64::kernel_ir_for(dt.to_dtype()); + let mut k = mt_moe_bgemm_iq2xxs_bm64::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; let gx = n_out / 64; let gy = m_total.div_ceil(64); @@ -191,7 +191,7 @@ fn run_gemv( buffers.insert("m_total".into(), (m_total as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("ctx"); - let mut k = ffai_moe_gemv_rows_iq2xxs::kernel_ir_for(dt.to_dtype()); + let mut k = mt_moe_gemv_rows_iq2xxs::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid(&k, &buffers, &BTreeMap::new(), [n_out, m_total, 1], [32, 1, 1]) diff --git a/crates/metaltile-std/tests/moe_gather_down_q2k_correctness.rs b/crates/metaltile-std/tests/moe_gather_down_q2k_correctness.rs index 4878cce5..1119f634 100644 --- a/crates/metaltile-std/tests/moe_gather_down_q2k_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_down_q2k_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::moe_gather_down_q2k` — fused 6-expert +//! GPU correctness for `kernels::moe::gather_down_q2k` — fused 6-expert //! Q2_K inline-dequant down-projection + router-weighted sum. Validates //! against a CPU reference running the identical (production-proven) //! Q2_K dequant formula. @@ -15,7 +15,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe_gather_down_q2k::ffai_moe_gather_down_q2k; +use metaltile_std::kernels::moe::gather_down_q2k::mt_moe_gather_down_q2k; // The Q2_K output-index → (qs byte, 2-bit shift) map is the single shared // definition in `quant::gguf`: the kernel, the quantizer, and this oracle all // read it, so the layout can't drift apart (getting it wrong was PR #264). @@ -111,7 +111,7 @@ fn run_gpu( buffers.insert("n_slots".into(), (N_SLOTS as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("Context::new on macOS"); - let mut kernel = ffai_moe_gather_down_q2k::kernel_ir_for(dt.to_dtype()); + let mut kernel = mt_moe_gather_down_q2k::kernel_ir_for(dt.to_dtype()); kernel.mode = KernelMode::Reduction; let result = ctx .dispatch_with_grid(&kernel, &buffers, &BTreeMap::new(), [m_out, 1, 1], [32, 1, 1]) diff --git a/crates/metaltile-std/tests/moe_gather_gemv_iq2xxs_correctness.rs b/crates/metaltile-std/tests/moe_gather_gemv_iq2xxs_correctness.rs index 7329edea..972ef2ea 100644 --- a/crates/metaltile-std/tests/moe_gather_gemv_iq2xxs_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_gemv_iq2xxs_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::moe_gather_gemv_iq2xxs` — the fused +//! GPU correctness for `kernels::moe::gather_gemv_iq2xxs` — the fused //! 6-expert IQ2_XXS inline-dequant gather GEMV used by the DSv4 decode //! FFN. Validates GPU output against a CPU reference that runs the //! identical (production-proven) IQ2_XXS dequant formula, so a wrong @@ -17,7 +17,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe_gather_gemv_iq2xxs::ffai_moe_gather_gemv_iq2xxs; +use metaltile_std::kernels::moe::gather_gemv_iq2xxs::mt_moe_gather_gemv_iq2xxs; const N_SLOTS: usize = 6; @@ -109,7 +109,7 @@ fn run_gpu( buffers.insert("m_out".into(), (m_out as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("Context::new on macOS"); - let mut kernel = ffai_moe_gather_gemv_iq2xxs::kernel_ir_for(dt.to_dtype()); + let mut kernel = mt_moe_gather_gemv_iq2xxs::kernel_ir_for(dt.to_dtype()); kernel.mode = KernelMode::Reduction; // grid (threadgroups) = [m_out, n_slots, 1], one 32-lane simdgroup each. diff --git a/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs index ca490c72..9cb0d3fc 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs @@ -24,7 +24,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe::{ +use metaltile_std::kernels::moe::orchestration::{ mt_moe_gather_qmm_int4_m8, mt_moe_gather_qmm_int4_m16, mt_moe_gather_qmm_int4_m32, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs b/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs index 2d9def0d..61deab89 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs @@ -13,7 +13,7 @@ use std::{collections::BTreeMap, time::Instant}; use common::{Dt, gpu_lock, pack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe::{ +use metaltile_std::kernels::moe::orchestration::{ mt_moe_gather_qmm_int4, mt_moe_gather_qmm_int4_m8, mt_moe_gather_qmm_mma_int4, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs index 67b81169..0fb80c12 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs @@ -3,7 +3,7 @@ #![allow(clippy::manual_is_multiple_of)] //! GPU correctness for the bit-width-generalized MMA MoE BGEMMs -//! `ffai::moe::mt_moe_gather_qmm_mma_b{3,5,6,8}`. +//! `kernels::moe::orchestration::mt_moe_gather_qmm_mma_b{3,5,6,8}`. //! //! Same tiled-MMA algorithm as `mt_moe_gather_qmm_mma_int4`, but the //! weight coop-dequant pulls codes from a contiguous LSB-first bit-stream @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe::{ +use metaltile_std::kernels::moe::orchestration::{ mt_moe_gather_qmm_mma_b3, mt_moe_gather_qmm_mma_b5, mt_moe_gather_qmm_mma_b6, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs index 975bc205..f755ccdc 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `ffai::moe_mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp`. +//! GPU correctness for `kernels::moe::mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp`. //! //! BM=BN=64 MPP MoE kernel — same output semantics as the BM=16 sibling but //! scaled up to a 64×64 output tile with 4 SGs (WM=WN=2) per TG. Validated @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{moe::mt_moe_gather_qmm_int4, moe_mpp_bm64}; +use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_int4, mpp_bm64}; /// Pack a row of int4 weights into uint32s (8 per uint, LSB-first per /// nibble). Identical to the helper used by the bm16_mpp test — @@ -148,7 +148,7 @@ fn moe_gather_qmm_mma_int4_bm64_mpp_matches_m1_clean_tile() { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - moe_mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); + mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [ceil(N/64), ceil(T/64), 1]. TG: 128 lanes = 4 SGs (WM=WN=2). let r = ctx @@ -267,7 +267,7 @@ fn moe_gather_qmm_mma_int4_bm64_mpp_matches_m1_multi_tile() { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - moe_mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); + mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid( @@ -391,7 +391,7 @@ fn moe_gather_qmm_mma_int4_bm64_mpp_bf16_matches_m1_clean_tile() { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - moe_mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::Bf16.to_dtype()); + mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::Bf16.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid( diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs index 51ac8065..997b433d 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `ffai::moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp`. +//! GPU correctness for `kernels::moe::mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp`. //! //! BM=BN=64 MPP MoE int8 kernel — same output semantics as the int4 BM=64 //! sibling but the weight layout changes from 8 nibbles/u32 to 4 bytes/u32. @@ -21,7 +21,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{moe::mt_moe_gather_qmm_b8, moe_mpp_bm64_int8}; +use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_b8, mpp_bm64_int8}; /// Pack a row of int8 weight codes into uint32s (4 codes per uint, LE byte /// order). Code values must be in 0..=255. @@ -134,7 +134,7 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_matches_b8_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( + let mut k = mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( Dt::F32.to_dtype(), ); k.mode = KernelMode::Reduction; @@ -250,7 +250,7 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_matches_b8_multi_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( + let mut k = mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( Dt::F32.to_dtype(), ); k.mode = KernelMode::Reduction; @@ -370,7 +370,7 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_bf16_matches_b8_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( + let mut k = mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( Dt::Bf16.to_dtype(), ); k.mode = KernelMode::Reduction; diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs index 7da304e8..163cc28a 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `ffai::moe_mpp_bm8::mt_moe_gather_qmm_mma_int4_bm8_mpp`. +//! GPU correctness for `kernels::moe::mpp_bm8::mt_moe_gather_qmm_mma_int4_bm8_mpp`. //! //! BM=8 MPP MoE kernel — same output semantics as the BM=16 / BM=64 siblings //! but the per-TG row tile shrinks to 8 to match decode-time MoE shapes @@ -23,7 +23,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{moe::mt_moe_gather_qmm_int4, moe_mpp_bm8}; +use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_int4, mpp_bm8}; /// Pack a row of int4 weights into uint32s (8 per uint, LSB-first per /// nibble). Same helper used by the bm16_mpp / bm64_mpp test files — @@ -187,7 +187,7 @@ fn run_case(case: &Case) { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm8::mt_moe_gather_qmm_mma_int4_bm8_mpp::kernel_ir_for(dt.to_dtype()); + let mut k = mpp_bm8::mt_moe_gather_qmm_mma_int4_bm8_mpp::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [ceil(N/32), ceil(T/8), 1]. TG: 32 lanes = 1 SG. let r = ctx diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs index 858951f3..8f6e99b2 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `ffai::moe_mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp`. +//! GPU correctness for `kernels::moe::mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp`. //! //! BM=8 MPP MoE int8 kernel — same output semantics as the int4 BM=8 sibling //! but the weight layout changes from 8 nibbles/u32 to 4 bytes/u32. @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{moe::mt_moe_gather_qmm_b8, moe_mpp_bm8_int8}; +use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_b8, mpp_bm8_int8}; /// Pack a row of int8 weight codes into uint32s (4 codes per uint, LE byte /// order). Code values must be in 0..=255. @@ -174,7 +174,7 @@ fn run_case(case: &Case) { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - moe_mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt.to_dtype()); + mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [ceil(N/32), ceil(T/8), 1]. TG: 32 lanes = 1 SG. let r = ctx diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs index 6b666c9b..f2ea2d26 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `ffai::moe_mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp`. +//! GPU correctness for `kernels::moe::mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp`. //! //! This is the MPP (MetalPerformancePrimitives) MoE BGEMM — same algorithm //! and output as `mt_moe_gather_qmm_mma_int4_bm16` but routes the inner @@ -23,7 +23,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{moe::mt_moe_gather_qmm_int4, moe_mpp}; +use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_int4, mpp}; /// Pack a row of int4 weights into uint32s (8 per uint, LSB-first per nibble). /// Identical to the helper used by the legacy @@ -141,7 +141,7 @@ fn moe_gather_qmm_mma_int4_bm16_mpp_matches_m1_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [N/BN=32, ceil(T/BM=16), 1]. TG: 32 lanes = 1 SG (MPP's // matmul2d uses `execution_simdgroup`). @@ -274,7 +274,7 @@ fn moe_gather_qmm_mma_int4_bm16_mpp_bf16_matches_m1_clean_tile() { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - moe_mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(Dt::Bf16.to_dtype()); + mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(Dt::Bf16.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid( diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs index 35772302..7341c920 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `ffai::moe_mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp`. +//! GPU correctness for `kernels::moe::mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp`. //! //! MPP (MetalPerformancePrimitives) int8 MoE BGEMM — same algorithm as //! `mt_moe_gather_qmm_mma_int4_bm16_mpp` but with pack-aligned 8-bit @@ -21,7 +21,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe_mpp_int8; +use metaltile_std::kernels::moe::mpp_int8; // ── helpers ──────────────────────────────────────────────────────────────── @@ -132,7 +132,7 @@ fn run_mpp_int8( buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("Context::new"); - let mut k = moe_mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt.to_dtype()); + let mut k = mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [N/BN=32, ceil(T/BM=16), 1], TG: [32, 1, 1] (1 SG — MPP matmul2d). let r = ctx diff --git a/crates/metaltile-std/tests/moe_gemv_rows_iq2xxs_correctness.rs b/crates/metaltile-std/tests/moe_gemv_rows_iq2xxs_correctness.rs index 9ecae4c6..4d55df72 100644 --- a/crates/metaltile-std/tests/moe_gemv_rows_iq2xxs_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_rows_iq2xxs_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_moe_gemv_rows_iq2xxs` — the fast +//! GPU correctness for `ffai::mt_moe_gemv_rows_iq2xxs` — the fast //! gemv-over-rows prefill MoE kernel (replaces the slow coop-tile bgemm). //! Oracle: per-row IQ2_XXS dequant gemv with PER-ROW x and per-row expert. //! Same dequant as the proven gather_gemv; only x is row-indexed. Cos ≥ 0.99. @@ -12,7 +12,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe_gemv_rows_iq2xxs::ffai_moe_gemv_rows_iq2xxs; +use metaltile_std::kernels::moe::gemv_rows_iq2xxs::mt_moe_gemv_rows_iq2xxs; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; @@ -94,7 +94,7 @@ fn gemv_rows_iq2xxs_matches_oracle() { buffers.insert("m_total".into(), (m_total as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = ffai_moe_gemv_rows_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = mt_moe_gemv_rows_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid(&k, &buffers, &BTreeMap::new(), [m_out, m_total, 1], [32, 1, 1]) diff --git a/crates/metaltile-std/tests/moe_gemv_rows_q2k_correctness.rs b/crates/metaltile-std/tests/moe_gemv_rows_q2k_correctness.rs index 2ecf06d3..3068fe8b 100644 --- a/crates/metaltile-std/tests/moe_gemv_rows_q2k_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_rows_q2k_correctness.rs @@ -1,7 +1,7 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_moe_gemv_rows_q2k` — proves it matches the -//! PROVEN pool bgemm `ffai_moe_gather_bgemm_q2k_mpp` on identical weights + +//! GPU correctness for `ffai::mt_moe_gemv_rows_q2k` — proves it matches the +//! PROVEN pool bgemm `mt_moe_gather_bgemm_q2k_mpp` on identical weights + //! per-row x (same canonical Q2_K dequant; only the GEMM structure differs). //! cosine ≥ 0.999. NO 86GB model load. #![cfg(target_os = "macos")] @@ -12,9 +12,9 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{ - moe_bgemm_q2k_mpp::ffai_moe_gather_bgemm_q2k_mpp, - moe_gemv_rows_q2k::ffai_moe_gemv_rows_q2k, +use metaltile_std::kernels::moe::{ + bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp, + gemv_rows_q2k::mt_moe_gemv_rows_q2k, }; fn xorshift(s: &mut u32) -> u32 { @@ -71,7 +71,7 @@ fn gemv_rows_q2k_matches_pool_kernel() { pool.insert("m_total".into(), (t_rows as u32).to_le_bytes().to_vec()); pool.insert("n_out".into(), (m_out as u32).to_le_bytes().to_vec()); pool.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); - let mut kp = ffai_moe_gather_bgemm_q2k_mpp::kernel_ir_for(Dt::F32.to_dtype()); + let mut kp = mt_moe_gather_bgemm_q2k_mpp::kernel_ir_for(Dt::F32.to_dtype()); kp.mode = KernelMode::Reduction; let rp = ctx .dispatch_with_grid(&kp, &pool, &BTreeMap::new(), [m_out / 32, t_rows.div_ceil(16), 1], [ @@ -92,7 +92,7 @@ fn gemv_rows_q2k_matches_pool_kernel() { gv.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); gv.insert("m_out".into(), (m_out as u32).to_le_bytes().to_vec()); gv.insert("m_total".into(), (t_rows as u32).to_le_bytes().to_vec()); - let mut kv = ffai_moe_gemv_rows_q2k::kernel_ir_for(Dt::F32.to_dtype()); + let mut kv = mt_moe_gemv_rows_q2k::kernel_ir_for(Dt::F32.to_dtype()); kv.mode = KernelMode::Reduction; let rv = ctx.dispatch_with_grid(&kv, &gv, &BTreeMap::new(), [m_out, t_rows, 1], [32, 1, 1]).unwrap(); diff --git a/crates/metaltile-std/tests/moe_gemv_rows_view_iq2xxs_correctness.rs b/crates/metaltile-std/tests/moe_gemv_rows_view_iq2xxs_correctness.rs index a8fc8f86..1b4b227d 100644 --- a/crates/metaltile-std/tests/moe_gemv_rows_view_iq2xxs_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_rows_view_iq2xxs_correctness.rs @@ -1,7 +1,7 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_moe_gemv_rows_view_iq2xxs` (u8-recombine -//! reads) and `ffai_moe_gemv_rows_view_u16_iq2xxs` (aligned u16 reads) — the +//! GPU correctness for `ffai::mt_moe_gemv_rows_view_iq2xxs` (u8-recombine +//! reads) and `mt_moe_gemv_rows_view_u16_iq2xxs` (aligned u16 reads) — the //! zero-copy gemv-over-rows MoE kernels that read raw 66-byte IQ2_XXS blocks //! straight from a no-copy view buffer. Oracle: per-row IQ2_XXS dequant gemv //! (same dequant as the proven `moe_gemv_rows_iq2xxs`), with the super-scale @@ -16,9 +16,9 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use half::f16; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe_gemv_rows_view_iq2xxs::{ - ffai_moe_gemv_rows_view_iq2xxs, - ffai_moe_gemv_rows_view_u16_iq2xxs, +use metaltile_std::kernels::moe::gemv_rows_view_iq2xxs::{ + mt_moe_gemv_rows_view_iq2xxs, + mt_moe_gemv_rows_view_u16_iq2xxs, }; fn xorshift(s: &mut u32) -> u32 { @@ -167,7 +167,7 @@ fn gemv_rows_view_iq2xxs_u8_matches_oracle() { .insert("expert_byte_stride".into(), (c.expert_byte_stride as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = ffai_moe_gemv_rows_view_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = mt_moe_gemv_rows_view_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid(&k, &buffers, &BTreeMap::new(), [c.m_out, c.m_total, 1], [32, 1, 1]) @@ -207,7 +207,7 @@ fn gemv_rows_view_u16_iq2xxs_matches_oracle() { .insert("expert_byte_stride".into(), (c.expert_byte_stride as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = ffai_moe_gemv_rows_view_u16_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = mt_moe_gemv_rows_view_u16_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid(&k, &buffers, &BTreeMap::new(), [c.m_out, c.m_total, 1], [32, 1, 1]) diff --git a/crates/metaltile-std/tests/moe_gemv_ws_iq2xxs_correctness.rs b/crates/metaltile-std/tests/moe_gemv_ws_iq2xxs_correctness.rs index 35868b3e..482283a5 100644 --- a/crates/metaltile-std/tests/moe_gemv_ws_iq2xxs_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_ws_iq2xxs_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_moe_gemv_ws_iq2xxs` — the WEIGHT-STATIONARY +//! GPU correctness for `ffai::mt_moe_gemv_ws_iq2xxs` — the WEIGHT-STATIONARY //! prefill MoE IQ2_XXS gemv (dequants each expert's weight row ONCE into //! threadgroup memory, reused across the tile's rows). Oracle: per-row //! IQ2_XXS dequant gemv from the SAME split pool (`qs_all`/`d_all`) as the @@ -14,7 +14,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::moe_gemv_ws_iq2xxs::ffai_moe_gemv_ws_iq2xxs; +use metaltile_std::kernels::moe::gemv_ws_iq2xxs::mt_moe_gemv_ws_iq2xxs; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; @@ -101,7 +101,7 @@ fn gemv_ws_iq2xxs_matches_oracle() { buffers.insert("rows_per_tile".into(), (rows_per_tile as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = ffai_moe_gemv_ws_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = mt_moe_gemv_ws_iq2xxs::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let gy = m_total.div_ceil(rows_per_tile); let r = diff --git a/crates/metaltile-std/tests/moe_gemv_ws_q2k_correctness.rs b/crates/metaltile-std/tests/moe_gemv_ws_q2k_correctness.rs index 34632228..e99f7057 100644 --- a/crates/metaltile-std/tests/moe_gemv_ws_q2k_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_ws_q2k_correctness.rs @@ -1,9 +1,9 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_moe_gemv_ws_q2k` — the WEIGHT-STATIONARY +//! GPU correctness for `ffai::mt_moe_gemv_ws_q2k` — the WEIGHT-STATIONARY //! prefill MoE Q2_K gemv (down projection). It dequants each expert's weight //! row ONCE into threadgroup memory and reuses it across the tile; the math is -//! identical to the proven `ffai_moe_gemv_rows_q2k` (same canonical Q2_K +//! identical to the proven `mt_moe_gemv_rows_q2k` (same canonical Q2_K //! dequant, same split pool, same per-row dot). Oracle = the gemv-rows kernel //! on the SAME inputs. Exact-ish f32 agreement (cosine ≥ 0.999), NO model load. #![cfg(target_os = "macos")] @@ -14,9 +14,9 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::ffai::{ - moe_gemv_rows_q2k::ffai_moe_gemv_rows_q2k, - moe_gemv_ws_q2k::ffai_moe_gemv_ws_q2k, +use metaltile_std::kernels::moe::{ + gemv_rows_q2k::mt_moe_gemv_rows_q2k, + gemv_ws_q2k::mt_moe_gemv_ws_q2k, }; fn xorshift(s: &mut u32) -> u32 { @@ -72,7 +72,7 @@ fn gemv_ws_q2k_matches_gemv_rows() { bref.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); bref.insert("m_out".into(), (m_out as u32).to_le_bytes().to_vec()); bref.insert("m_total".into(), (m_total as u32).to_le_bytes().to_vec()); - let mut kr = ffai_moe_gemv_rows_q2k::kernel_ir_for(Dt::F32.to_dtype()); + let mut kr = mt_moe_gemv_rows_q2k::kernel_ir_for(Dt::F32.to_dtype()); kr.mode = KernelMode::Reduction; let rr = ctx .dispatch_with_grid(&kr, &bref, &BTreeMap::new(), [m_out, m_total, 1], [32, 1, 1]) @@ -92,7 +92,7 @@ fn gemv_ws_q2k_matches_gemv_rows() { bws.insert("m_out".into(), (m_out as u32).to_le_bytes().to_vec()); bws.insert("m_total".into(), (m_total as u32).to_le_bytes().to_vec()); bws.insert("rows_per_tile".into(), (rows_per_tile as u32).to_le_bytes().to_vec()); - let mut kw = ffai_moe_gemv_ws_q2k::kernel_ir_for(Dt::F32.to_dtype()); + let mut kw = mt_moe_gemv_ws_q2k::kernel_ir_for(Dt::F32.to_dtype()); kw.mode = KernelMode::Reduction; let gy = m_total.div_ceil(rows_per_tile); let rw = diff --git a/crates/metaltile-std/tests/moe_q2k_view_u16_correctness.rs b/crates/metaltile-std/tests/moe_q2k_view_u16_correctness.rs index 03fa1b14..085ec2bf 100644 --- a/crates/metaltile-std/tests/moe_q2k_view_u16_correctness.rs +++ b/crates/metaltile-std/tests/moe_q2k_view_u16_correctness.rs @@ -12,9 +12,9 @@ use metaltile::{ Context, core::{dtype::DType, ir::KernelMode}, }; -use metaltile_std::ffai::{ - moe_bgemm_q2k_bm64::ffai_moe_bgemm_q2k_bm64, - moe_bgemm_q2k_view_u16_bm64::ffai_moe_bgemm_q2k_view_u16_bm64, +use metaltile_std::kernels::moe::{ + bgemm_q2k_bm64::mt_moe_bgemm_q2k_bm64, + bgemm_q2k_view_u16_bm64::mt_moe_bgemm_q2k_view_u16_bm64, }; fn xs(s: &mut u32) -> u32 { @@ -106,7 +106,7 @@ fn q2k_view_u16_bm64_matches_pool() { pb.insert("m_total".into(), (m_total as u32).to_le_bytes().to_vec()); pb.insert("n_out".into(), (n_out as u32).to_le_bytes().to_vec()); pb.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); - let pool_out = run(pb, ffai_moe_bgemm_q2k_bm64::kernel_ir_for(DType::F32)); + let pool_out = run(pb, mt_moe_bgemm_q2k_bm64::kernel_ir_for(DType::F32)); let raw_u16: Vec = raw.clone(); // bytes; view_u16/view_f16 reinterpret let mut vb: BTreeMap> = BTreeMap::new(); @@ -120,7 +120,7 @@ fn q2k_view_u16_bm64_matches_pool() { vb.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); vb.insert("tensor_byte_off".into(), 0u32.to_le_bytes().to_vec()); vb.insert("expert_byte_stride".into(), ((nblk * 84) as u32).to_le_bytes().to_vec()); - let view_out = run(vb, ffai_moe_bgemm_q2k_view_u16_bm64::kernel_ir_for(DType::F32)); + let view_out = run(vb, mt_moe_bgemm_q2k_view_u16_bm64::kernel_ir_for(DType::F32)); let mut worst = 0.0f32; let mut wi = 0; diff --git a/crates/metaltile-std/tests/moe_view_u16_correctness.rs b/crates/metaltile-std/tests/moe_view_u16_correctness.rs index 4bca114f..3fc7b4a7 100644 --- a/crates/metaltile-std/tests/moe_view_u16_correctness.rs +++ b/crates/metaltile-std/tests/moe_view_u16_correctness.rs @@ -16,9 +16,9 @@ use metaltile::{ Context, core::{dtype::DType, ir::KernelMode}, }; -use metaltile_std::ffai::{ - moe_bgemm_iq2xxs_bm64::ffai_moe_bgemm_iq2xxs_bm64, - moe_bgemm_iq2xxs_view_u16_bm64::ffai_moe_bgemm_iq2xxs_view_u16_bm64, +use metaltile_std::kernels::moe::{ + bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, + bgemm_iq2xxs_view_u16_bm64::mt_moe_bgemm_iq2xxs_view_u16_bm64, }; struct Lcg(u64); @@ -152,7 +152,7 @@ fn view_u16_bm64_matches_pool_bm64() { pb.insert("m_total".into(), (m_total as u32).to_le_bytes().to_vec()); pb.insert("n_out".into(), (n_out as u32).to_le_bytes().to_vec()); pb.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); - let pool_out = run(pb, ffai_moe_bgemm_iq2xxs_bm64::kernel_ir_for(DType::F32)); + let pool_out = run(pb, mt_moe_bgemm_iq2xxs_bm64::kernel_ir_for(DType::F32)); // VIEW-u16 bm64 let raw_bytes: Vec = raw_u16.iter().flat_map(|v| v.to_le_bytes()).collect(); @@ -169,7 +169,7 @@ fn view_u16_bm64_matches_pool_bm64() { vb.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); vb.insert("tensor_byte_off".into(), 0u32.to_le_bytes().to_vec()); vb.insert("expert_byte_stride".into(), ((nblk * 66) as u32).to_le_bytes().to_vec()); - let view_out = run(vb, ffai_moe_bgemm_iq2xxs_view_u16_bm64::kernel_ir_for(DType::F32)); + let view_out = run(vb, mt_moe_bgemm_iq2xxs_view_u16_bm64::kernel_ir_for(DType::F32)); let mut worst = 0.0f32; let mut wi = 0; diff --git a/docs/specs/KERNEL_AUDIT.md b/docs/specs/KERNEL_AUDIT.md index a24de091..e8d349d8 100644 --- a/docs/specs/KERNEL_AUDIT.md +++ b/docs/specs/KERNEL_AUDIT.md @@ -104,8 +104,8 @@ Local verification of NAX kernels is the developer's responsibility on M4+ hardw | gemv_masked | ✓ | ✓ | ✓ | `kernels/gemm/gemv_masked.rs` → `mt_gemv_masked`. | | quantized (affine_quantize / affine_dequantize) | ✓ | ✓ | ✓ | `kernels/gemm/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) | ✓ | ✓ | ✓ | `kernels/gemm/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. | -| moe (router top-k + permute + unpermute orchestration) | ✗ | ✓ | ✓ | `ffai/moe.rs` → `mt_moe_router_topk`, `mt_moe_permute`, `mt_moe_unpermute`. MoE expert-routing orchestration. The grouped quantized BGEMM that fuses per-expert FFN matmuls is counted under the `quantized (gather_*)` row. | +| quantized (gather_qmv / gather_qmm — gather variants) | ✓ | ✓ | ✓ | `kernels/moe/orchestration.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` (`kernels/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): `mt_dequant_gemv_int4_expert_indexed` — per-output-row expert selection for the gate/up FFN dispatch shape. | +| moe (router top-k + permute + unpermute orchestration) | ✗ | ✓ | ✓ | `kernels/moe/orchestration.rs` → `mt_moe_router_topk`, `mt_moe_permute`, `mt_moe_unpermute`. MoE expert-routing orchestration. The grouped quantized BGEMM that fuses per-expert FFN matmuls is counted under the `quantized (gather_*)` row. | | dequant_gather (quantized embedding-table gather) | ✗ | ✗ | ✓ | `ffai/dequant_gather.rs`. int{3,4,5,6,8} all bit-widths. FFAI-only. | | dequant_gemv (quantized GEMV, FFAI flavour) | ~ | ~ | ✓ | `kernels/gemm/dequant_gemv.rs` → `mt_dequant_gemv_int{2,3,4,5,6,8}` (one-row-per-TG) + `mt_dequant_gemv_int4_fast` (PR #154, 8-row-per-TG, mirrors MLX `qmv_fast`). The non-fast int4 kernel stays because FFAI's GPU-router opts into its indirect Swift wrapper. | | fp_quantized (fp4/fp8 quant + dequant) | ✓ | ✓ | ✓ | `kernels/gemm/fp_quantized.rs` → `mt_fp4_quant_dequant` (fp4 E2M1) + `mt_fp8_e4m3_quant_dequant` / `mt_fp8_e5m2_quant_dequant` (fp8). Pure arithmetic transform (per-group max-scale + mantissa rounding via `floor(log2)`/`exp2`/`round`); exact for fp8 normals/subnormals, saturating (no NaN/Inf). | @@ -251,9 +251,9 @@ each family's proven dispatch geometry verbatim (no new freeze surface). | qmm — simdgroup-MMA | simdgroup-matrix | `kernels/gemm/block_scaled_mma.rs` | int4, int8 | | qmm — MPP (tensor engine) | MPP `matmul2d` | `kernels/gemm/block_scaled_qmm_mpp.rs` | int4, int8 | | qmm — NAX | NAX `matmul2d` | `kernels/gemm/block_scaled_qmm_nax.rs` | int4, int8 | -| MoE gather-qmm | reduction | `mlx/block_scaled_moe.rs` | int3–8 | -| MoE gather — MPP (bm8/16/64) | MPP | `ffai/moe_mpp{,_bm8,_bm64}_block_scaled.rs` | int4, int8 | -| expert-indexed GEMV | reduction | `ffai/dequant_gemv_expert_indexed_block_scaled.rs` | int4 | +| MoE gather-qmm | reduction | `kernels/moe/block_scaled.rs` | int3–8 | +| MoE gather — MPP (bm8/16/64) | MPP | `kernels/moe/mpp{,_bm8,_bm64}_block_scaled.rs` | int4, int8 | +| expert-indexed GEMV | reduction | `kernels/moe/dequant_gemv_expert_indexed_block_scaled.rs` | int4 | | fused RMSNorm + GEMV | reduction | `kernels/norm/rms_norm_block_scaled_qgemv.rs` | int4, int8-fast | | fused gated-RMSNorm + GEMV | reduction | `kernels/norm/gated_rms_norm_block_scaled_qgemv.rs` | int4 | | batched-Q/K/V qgemv + qmm | reduction | `kernels/gemm/batched_qkv_block_scaled_{qgemv,qmm}.rs` | int4, int8-fast | @@ -335,7 +335,7 @@ A few rows mix multiple `.metal` files into one op or split one file into multip - **`steel/`** — each kernel file becomes one op row; per-block-shape instantiations are not counted separately. `steel_attention` (scalar) and `steel_attention_mma` (simdgroup-MMA) are two rows because they are separately compiled kernels with different lowering strategies. - **`quantized.metal`** — split into four rows by semantic operation (quant/dequant, qmv/qvm/qmm matmul, gather-qmv/qmm, fp4/fp8). The Apple10+ variants (`quantized_nax`, `fp_quantized_nax`) are separate rows because they live in separate modules with runtime-only dispatch gating. `fp_quantized_mma` is its own row (runs on M1+, no Apple10 gating). - **`indexing/`** is one row covering scatter / scatter_axis / gather_axis / gather_front / masked_scatter. Bare `gather` is its own row (FFAI-specific). -- **`moe`** is the routing/permute/unpermute orchestration in `ffai/moe.rs`. The grouped quantized BGEMM lives under the `quantized (gather_*)` row. +- **`moe`** is the routing/permute/unpermute orchestration in `kernels/moe/orchestration.rs`. The grouped quantized BGEMM lives under the `quantized (gather_*)` row. - **`logits processors`** is one row for the FFAI sampler-stage kernels (`temperature`, `repetition_penalty`, `topk` / `top_p` / `min_p` masks). - Cells marked **`~`** indicate a partial port (typically one bit-width, one dtype, or one block shape where upstream has many) — see the notes column for the specific gap. diff --git a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md index 62dbdb79..9a64f37e 100644 --- a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md +++ b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md @@ -55,9 +55,12 @@ crates/metaltile-std/src/kernels/ 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 - moe/ 🔨 SEEDED (folder created early) — gather_q4 (batched expert up/down/weighted-sum) - · sigmoid_bias (router pre-score), split out of the gemv_q8 grab-bag. Remaining: - moe orchestration · mpp(bm8/bm64 × int8) · bgemm/gemv(q2k/iq2xxs) · block_scaled_moe + moe/ ✅ DONE — orchestration (router_topk + permute/unpermute + gather_qmm) · + router_topk_biased / sigmoid_bias / sqrtsoftplus · mpp(bm8/bm64 × int8 × + block_scaled) + mpp_shared · bgemm/gemv (q2k/iq2xxs/q4, view/ws/rows) · gather_q4 · + down_swiglu_accum / down_weighted_sum · dequant_gemv_expert_indexed(_block_scaled) · + block_scaled. Filenames drop the moe_ prefix; format-axis fold deferred (§7). + orchestration.rs (~4k lines) slated for a follow-up split. norm/ ✅ DONE — rms_norm(+residual/rope/qgemv/gated) · layer_norm · adain1d rope/ ✅ DONE — rope · rope_2d · rope_banded · rope_yarn · partial_rope convolution/ ✅ DONE — conv1d/2d/3d · depthwise · winograd · steel_conv · conv1d_causal(_roll) (see §4) @@ -166,7 +169,7 @@ payoff last: |---|---|---|---| | ✅ done | `convolution/`, `rope/`, `norm/`, `sampling/`, `ops/` | exemplar + all of wave 1 | 24k → ~1.6k | | 2 | ✅ `audio/` `vision/` `kv_cache/` `gemm/` (dense + quantized) `ssm/` all done | 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 | +| 3 | ✅ `moe/` done; remaining `sdpa/`, **`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 From 5af6db0c72486424b2744aeb5e10824f2f12ba12 Mon Sep 17 00:00:00 2001 From: Eric Kryski <599019+ekryski@users.noreply.github.com> Date: Sat, 13 Jun 2026 13:31:44 -0600 Subject: [PATCH 2/7] refactor(moe): split orchestration.rs into router_topk / permute / gather_qmm MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The former moe.rs (~4.3k lines, 13 kernels + a shared test/bench module) was moved whole in the previous commit; split it now per the <1000-line guideline: - router_topk.rs — mt_moe_router_topk (top-k expert selection) - permute.rs — mt_moe_permute / mt_moe_unpermute - gather_qmm.rs — the 10 grouped-gather quantized-matmul cells (int4/int8/ b{3,5,6,8}, m8/m16/m32, mma, bm16) + their CSR/MMA helpers Each file carries its own kernel_tests + kernel_benches (the shared u32_bytes helper is duplicated into router_topk/permute). Consumer test imports remapped orchestration:: -> the new modules. Kernel set unchanged (1272 codegen). --- .../moe/{orchestration.rs => gather_qmm.rs} | 503 +----------------- crates/metaltile-std/src/kernels/moe/mod.rs | 8 +- .../metaltile-std/src/kernels/moe/permute.rs | 240 +++++++++ .../src/kernels/moe/router_topk.rs | 274 ++++++++++ ...moe_gather_qmm_int4_m16_m32_correctness.rs | 2 +- .../tests/moe_gather_qmm_microbench.rs | 2 +- ...moe_gather_qmm_mma_bitwidth_correctness.rs | 4 +- .../moe_gather_qmm_mpp_bm64_correctness.rs | 2 +- ...oe_gather_qmm_mpp_bm64_int8_correctness.rs | 2 +- .../moe_gather_qmm_mpp_bm8_correctness.rs | 2 +- ...moe_gather_qmm_mpp_bm8_int8_correctness.rs | 2 +- .../tests/moe_gather_qmm_mpp_correctness.rs | 2 +- docs/specs/KERNEL_AUDIT.md | 6 +- docs/specs/KERNEL_CONSOLIDATION_PLAN.md | 4 +- 14 files changed, 538 insertions(+), 515 deletions(-) rename crates/metaltile-std/src/kernels/moe/{orchestration.rs => gather_qmm.rs} (87%) create mode 100644 crates/metaltile-std/src/kernels/moe/permute.rs create mode 100644 crates/metaltile-std/src/kernels/moe/router_topk.rs diff --git a/crates/metaltile-std/src/kernels/moe/orchestration.rs b/crates/metaltile-std/src/kernels/moe/gather_qmm.rs similarity index 87% rename from crates/metaltile-std/src/kernels/moe/orchestration.rs rename to crates/metaltile-std/src/kernels/moe/gather_qmm.rs index 56322bbe..8d723ab3 100644 --- a/crates/metaltile-std/src/kernels/moe/orchestration.rs +++ b/crates/metaltile-std/src/kernels/moe/gather_qmm.rs @@ -1,306 +1,13 @@ //! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric //! SPDX-License-Identifier: Apache-2.0 -//! MoE orchestration kernels — router top-k, permute, unpermute, -//! grouped BGEMM dispatch. -//! -//! Targets Qwen3.6-35B-A3B and Qwen3-Coder-30B-A3B end-to-end serving. -//! The per-expert quantized matmul cell is already served by -//! `mt_qmm_*` (mma / mma_m16 / bm4 / bm2 / v2) — this module adds the -//! routing kernels that go around each expert call. -//! -//! ## Pipeline shape -//! -//! ```text -//! activations [B*T, hidden] -//! │ -//! ▼ -//! ┌──────────────────┐ -//! │ mt_moe_router_topk│ logits → [B*T, k] (indices + weights) -//! └──────────────────┘ -//! │ -//! ▼ -//! ┌──────────────────┐ -//! │ mt_moe_permute │ [B*T, hidden] → [k*B*T, hidden] expert-sorted -//! └──────────────────┘ -//! │ -//! ▼ -//! ┌──────────────────┐ -//! │ per-expert qmm │ N × mt_qmm_for() calls — already shipped -//! └──────────────────┘ -//! │ -//! ▼ -//! ┌──────────────────┐ -//! │ mt_moe_unpermute │ [k*B*T, hidden] + weights → [B*T, hidden] -//! └──────────────────┘ -//! ``` +//! MoE grouped-gather quantized matmul — the per-expert weight×activation +//! cell that runs inside each routed expert call. CSR-routed scalar forms +//! (int4 + b{3,5,6,8}), the m8/m16/m32 short-prefill tilings, and the +//! simdgroup-MMA forms (int4 / int8 / b{3,5,6,8}, bm16). Per-row `indices` +//! or CSR `expert_offsets` select each token's expert. use metaltile::kernel; -// ── mt_moe_router_topk ─────────────────────────────────────────────────── -// -// Per-token select top-k experts from `router_logits`, plus softmax -// weights over the chosen k. -// -// Inputs: -// router_logits — [B*T, n_experts] (any float dtype, computed in f32) -// indices_out — [B*T, k] (u32) -// weights_out — [B*T, k] (same dtype as router_logits, softmax weights) -// -// Constexpr: -// n_experts — typical Qwen3.6-A3B: 128. Must fit one simdgroup -// (≤ 32×32 = 1024) — every reasonable MoE topology. -// k — typical 6-8 for production MoE. Hard cap k ≤ 32. -// -// Geometry: -// tpg=32 (one simdgroup per token row) -// grid = [B*T, 1, 1] (Reduction mode) -// -// Algorithm — k iterations of simd-parallel argmax with mask of -// previously-chosen indices stored in TG memory. After k passes, -// softmax over the chosen k values in-place on lane 0..k-1. -// -// Bench spec uses BenchDispatch::Generic + shapes: &[] so `tile bench` -// skips it; correctness lives in unit tests + downstream MoE -// integration. Same convention as other ffai/ kernels (gather, sampling). -#[kernel] -pub fn mt_moe_router_topk( - router_logits: Tensor, - mut indices_out: Tensor, - mut weights_out: Tensor, - #[constexpr] n_experts: u32, - #[constexpr] k: u32, - // 1 = Qwen3-MoE style (softmax over chosen-k, sum-to-1 — `norm_topk_prob=True`) - // 0 = Qwen3-Next style (softmax over ALL n_experts, return chosen probs - // un-renormalized — `norm_topk_prob=False`) - // Mathematically equivalent at mode 1: softmax-over-chosen-k is the - // same as (softmax-over-all → renormalize-over-chosen). Mode 0 - // returns probs that sum to < 1 across the chosen k, matching MLX's - // qwen3_next.py:334-341. - // - // INVARIANT: this kernel pins tpg=32 (one simdgroup per token row). - // The `simdgroup_barrier_mem_none()` below is correct only at tpg=32. - // Caller must dispatch with `[n_rows, 1, 1] × [32, 1, 1]`. - #[constexpr] norm_topk_prob: u32, -) { - let row = tgid_x; - let lane = tid; - let row_base = row * n_experts; - // TG scratch: chosen indices + values from each of the k argmax passes. - // 32 slots covers any reasonable k (typical 6-8). Kernel assumes - // k ≤ 32 — caller MUST enforce this in the host-side dispatcher - // (no GPU-side check, would silently scribble into adjacent TG mem). - threadgroup_alloc("tg_chosen_idx", 32u32); - threadgroup_alloc("tg_chosen_val", 32u32); - // Cache the all-experts-softmax sum for Qwen3-Next mode (mode 0). - // 1 slot, written by lane 0 in the prepass. - threadgroup_alloc("tg_full_sum", 1u32); - threadgroup_alloc("tg_full_max", 1u32); - // ── Pre-pass: compute softmax denominator over ALL n_experts ───── - // Needed only for norm_topk_prob=0 (Qwen3-Next), but the cost is - // trivial (one simd_max + simd_sum) and emitting it unconditionally - // keeps the codegen tight (the codegen DCE will drop the dead path - // when the constexpr branch is unreachable). - let mut local_max_all = neg_infinity(); - let n_per_lane_pre = (n_experts + 31u32) / 32u32; - for r in range(0u32, n_per_lane_pre, 1u32) { - let j = r * 32u32 + lane; - if j < n_experts { - let v = load(router_logits[row_base + j]).cast::(); - let better = v > local_max_all; - local_max_all = select(better, v, local_max_all); - } - } - let row_max_all = simd_max(local_max_all); - let mut local_sum_all = 0.0f32; - for r in range(0u32, n_per_lane_pre, 1u32) { - let j = r * 32u32 + lane; - if j < n_experts { - let v = load(router_logits[row_base + j]).cast::(); - local_sum_all = local_sum_all + exp(v - row_max_all); - } - } - let row_sum_all = simd_sum(local_sum_all); - if lane == 0u32 { - threadgroup_store("tg_full_max", 0u32, row_max_all); - threadgroup_store("tg_full_sum", 0u32, row_sum_all); - } - simdgroup_barrier_mem_none(); - // ── k argmax passes with chosen-mask ───────────────────────────── - for it in range(0u32, k, 1u32) { - // Per-lane local argmax over its slice of n_experts. - // Each lane covers ceil(n_experts/32) experts. - let mut best_val = neg_infinity(); - let mut best_idx = 0u32; - let n_per_lane = (n_experts + 31u32) / 32u32; - for r in range(0u32, n_per_lane, 1u32) { - let j = r * 32u32 + lane; - if j < n_experts { - let v = load(router_logits[row_base + j]).cast::(); - // Mask: was j picked in a previous iter? - // Scan tg_chosen_idx[0..it] — k ≤ 8 typically so this - // is fast even without early exit. - let mut chosen_mask = 0u32; - for p in range(0u32, it, 1u32) { - let cp = threadgroup_load("tg_chosen_idx", p); - chosen_mask = chosen_mask | select(j == cp, 1u32, 0u32); - } - let candidate = select(chosen_mask > 0u32, neg_infinity(), v); - let better = candidate > best_val; - best_val = select(better, candidate, best_val); - best_idx = select(better, j, best_idx); - } - } - // Cross-lane reduce. simd_max gives the global best value; - // ties broken to smaller idx via simd_min on (idx | sentinel). - let global_best_val = simd_max(best_val); - let i_have = best_val == global_best_val; - let my_idx_or_max = select(i_have, best_idx, 4294967295u32); // u32::MAX - let global_best_idx = simd_min(my_idx_or_max); - // Lane 0 writes the iter's chosen slot. - if lane == 0u32 { - threadgroup_store("tg_chosen_idx", it, global_best_idx); - threadgroup_store("tg_chosen_val", it, global_best_val); - } - simdgroup_barrier_mem_none(); - } - // ── Softmax / weight emit per `norm_topk_prob` ────────────────── - // Mode 1 (Qwen3-MoE, default): softmax over chosen-k (sum-to-1). - // numerator = exp(z_i - max_chosen); divisor = Σ_j∈chosen - // == exp(z_i - max_all) · const / Σ_j∈chosen exp(z_j - max_all) · const - // so we can use the SAME numerator as mode 0 (exp(z - max_all)) and - // just swap the divisor. Avoids needing a Rust `if`-expression - // which the DSL doesn't unify across arms. - // Mode 0 (Qwen3-Next): un-normalized chosen probs (sum < 1). - // weight_i = exp(z_i - max_all) / Σ_j∈all exp(z_j - max_all) - let my_val = select(lane < k, threadgroup_load("tg_chosen_val", lane), neg_infinity()); - let row_max_full = threadgroup_load("tg_full_max", 0u32); - let row_sum_full = threadgroup_load("tg_full_sum", 0u32); - let exp_val = exp(my_val - row_max_full); - let masked_exp = select(lane < k, exp_val, 0.0f32); - let sum_chosen = simd_sum(masked_exp); - // Pick divisor: chosen-k sum for renormalized (mode 1) or all-experts - // sum for raw probs (mode 0). select() forces both to be live; codegen - // const-folds when `norm_topk_prob` bakes in. - let divisor = select(norm_topk_prob == 1u32, sum_chosen, row_sum_full); - let weight = masked_exp / divisor; - // ── Write outputs ─────────────────────────────────────────────── - if lane < k { - let out_base = row * k + lane; - store(indices_out[out_base], threadgroup_load("tg_chosen_idx", lane)); - store(weights_out[out_base], weight.cast::()); - } -} - -// ── mt_moe_unpermute ───────────────────────────────────────────────────── -// -// Combine k expert outputs back into the original token order with -// top-k softmax weights. -// -// Inputs: -// expert_outputs — [k*B*T, hidden] per-expert dense outputs at the -// expert-sorted positions -// inv_perm — [B*T, k] where (token i, slot j) was placed -// in expert_outputs (computed by -// caller's sort step) -// top_k_weights — [B*T, k] softmax weights from -// mt_moe_router_topk -// out — [B*T, hidden] weighted sum across k experts -// -// Constexpr: -// hidden — model hidden dim (e.g. 2048 for Qwen3-MoE) -// k — top-k expert count (e.g. 8) -// -// Geometry: -// tpg=128 (split hidden across 128 lanes via 4-wide vectorize) -// grid=[B*T, 1, 1] -// -// Per-token cost: read k * hidden / 128 = (k * hidden) / 128 expert -// values + k weights, do k FMAs per output column, one store per -// column. At hidden=2048, k=8 → ~1k FMAs per token. Bandwidth-bound, -// not ALU-bound. -#[kernel] -pub fn mt_moe_unpermute( - expert_outputs: Tensor, - inv_perm: Tensor, - top_k_weights: Tensor, - mut out: Tensor, - #[constexpr] hidden: u32, - #[constexpr] k: u32, -) { - let token = tgid_x; - let lane = tid; - let row_base_inv = token * k; - let row_base_w = token * k; - let row_base_out = token * hidden; - let n_per_lane = (hidden + 127u32) / 128u32; - for r in range(0u32, n_per_lane, 1u32) { - let h = r * 128u32 + lane; - if h < hidden { - let mut acc = 0.0f32; - for j in range(0u32, k, 1u32) { - let pos = load(inv_perm[row_base_inv + j]); - let v = load(expert_outputs[pos * hidden + h]).cast::(); - let w = load(top_k_weights[row_base_w + j]).cast::(); - acc = acc + w * v; - } - store(out[row_base_out + h], acc.cast::()); - } - } -} - -// ── mt_moe_permute ─────────────────────────────────────────────────────── -// -// Gather tokens into per-expert contiguous buffers given a pre-computed -// sort permutation. The expensive sort step (argsort over top-k expert -// indices) is done by the caller — typically CPU-side via Rust sort, -// or via a future sort kernel. This kernel is just the data-movement -// half: each output position copies the row indicated by sort_token_idx. -// -// Inputs: -// tokens — [B*T, hidden] activations to gather -// sort_token_idx — [k * B*T] for each permuted position p, -// which original token row sourced it. -// Caller computes via argsort over -// top-k indices flattened to -// (token * k + slot) → token (this is -// the "permute" direction; the inverse -// is `inv_perm` consumed by unpermute). -// permuted — [k * B*T, hidden] expert-sorted output. Each k*B*T -// row corresponds to one (expert, token) -// pair; consecutive rows with the same -// expert form that expert's input slab. -// -// Constexpr: -// hidden — model hidden dim -// -// Geometry: -// tpg=128 (split hidden across 128 lanes, ceil(hidden/128) iters) -// grid=[k*B*T, 1, 1] -// -// Per-permuted-row cost: hidden / 128 = 16 loads + 16 stores (at -// hidden=2048). Bandwidth-bound — no FMAs, just a vector copy. -#[kernel] -pub fn mt_moe_permute( - tokens: Tensor, - sort_token_idx: Tensor, - mut permuted: Tensor, - #[constexpr] hidden: u32, -) { - let permuted_pos = tgid_x; - let lane = tid; - let token = load(sort_token_idx[permuted_pos]); - let src_base = token * hidden; - let dst_base = permuted_pos * hidden; - let n_per_lane = (hidden + 127u32) / 128u32; - for r in range(0u32, n_per_lane, 1u32) { - let h = r * 128u32 + lane; - if h < hidden { - let v = load(tokens[src_base + h]); - store(permuted[dst_base + h], v); - } - } -} - // ── mt_moe_gather_qmm_int4 ──────────────────────────────────────────────── // // Grouped quantized matmul for MoE. Matches MLX's `gatherQuantizedMM` @@ -3910,137 +3617,8 @@ pub mod kernel_tests { fn test_moe_gather_qmm_mma_int8(dt: DType) -> TestSetup { mma_setup(mt_moe_gather_qmm_mma_int8::kernel_ir_for(dt), 8, 32, 128, dt) } - - // ── Router top-k + permute / unpermute orchestration ────────────────── - - /// Router oracle: softmax over all experts for the denominator, pick the - /// top-k logits (well-separated test inputs → no ties), then weight either - /// by renormalised softmax over the chosen k (`norm_topk_prob`) or by the - /// global softmax (raw probs that sum to < 1). - fn router_oracle( - logits: &[f32], - n_rows: usize, - n_experts: usize, - k: usize, - norm_topk_prob: bool, - ) -> (Vec, Vec) { - let mut idx_out = vec![0u32; n_rows * k]; - let mut w_out = vec![0.0f32; n_rows * k]; - for row in 0..n_rows { - let row_l = &logits[row * n_experts..(row + 1) * n_experts]; - let max_all = row_l.iter().copied().fold(f32::NEG_INFINITY, f32::max); - let sum_all: f32 = row_l.iter().map(|&l| (l - max_all).exp()).sum(); - // Top-k by descending logit (stable: smaller index wins ties). - let mut order: Vec = (0..n_experts).collect(); - order.sort_by(|&a, &b| row_l[b].partial_cmp(&row_l[a]).unwrap().then(a.cmp(&b))); - let chosen = &order[..k]; - let sum_chosen: f32 = chosen.iter().map(|&e| (row_l[e] - max_all).exp()).sum(); - for (i, &e) in chosen.iter().enumerate() { - idx_out[row * k + i] = e as u32; - let num = (row_l[e] - max_all).exp(); - w_out[row * k + i] = if norm_topk_prob { num / sum_chosen } else { num / sum_all }; - } - } - (idx_out, w_out) - } - - fn router_setup(dt: DType, norm_topk_prob: bool) -> TestSetup { - let (n_rows, n_experts, k) = (4usize, 8usize, 4usize); - // Well-separated logits (distinct multiples of 0.5 per row → no ties, - // gap ≫ dtype epsilon so the selection is dtype-stable). - let logits_f: Vec = (0..n_rows * n_experts) - .map(|i| { - let row = i / n_experts; - let e = i % n_experts; - ((e * 5 + row * 3) % n_experts) as f32 * 0.5 - }) - .collect(); - let logits = unpack_f32(&pack_f32(&logits_f, dt), dt); - let (idx, w) = router_oracle(&logits, n_rows, n_experts, k, norm_topk_prob); - TestSetup::new(mt_moe_router_topk::kernel_ir_for(dt)) - .mode(KernelMode::Reduction) - .input(TestBuffer::from_vec("router_logits", pack_f32(&logits_f, dt), dt)) - .input(TestBuffer::zeros("indices_out", n_rows * k, DType::U32)) - .input(TestBuffer::zeros("weights_out", n_rows * k, dt)) - .constexpr("n_experts", n_experts as u32) - .constexpr("k", k as u32) - .constexpr("norm_topk_prob", u32::from(norm_topk_prob)) - .expect(TestBuffer::from_vec("indices_out", u32_bytes(&idx), DType::U32)) - .expect(TestBuffer::from_vec("weights_out", pack_f32(&w, dt), dt)) - .grid_3d(n_rows as u32, 1, 1, [32, 1, 1]) - } - - // norm_topk_prob = 1: weights renormalised over the chosen k (Qwen3-MoE). - #[test_kernel(dtypes = [f32, f16, bf16], tol = [1e-4, 1e-2, 5e-2])] - fn test_moe_router_topk_norm(dt: DType) -> TestSetup { router_setup(dt, true) } - // norm_topk_prob = 0: raw global-softmax probs (Qwen3-Next). - #[test_kernel(dtypes = [f32, f16, bf16], tol = [1e-4, 1e-2, 5e-2])] - fn test_moe_router_topk_global(dt: DType) -> TestSetup { router_setup(dt, false) } - - // mt_moe_permute: gather token rows into expert-sorted order. - #[test_kernel(dtypes = [f32, f16, bf16], tol = [1e-4, 1e-3, 1e-2])] - fn test_moe_permute(dt: DType) -> TestSetup { - let (n_tokens, k, hidden) = (4usize, 2usize, 128usize); - let n_permuted = k * n_tokens; - let tokens_f: Vec = - (0..n_tokens * hidden).map(|i| ((i as f32) * 0.013).sin()).collect(); - let sort_token_idx: Vec = vec![0, 2, 1, 3, 1, 0, 3, 2]; - let tokens = unpack_f32(&pack_f32(&tokens_f, dt), dt); - let mut expected = vec![0.0f32; n_permuted * hidden]; - for p in 0..n_permuted { - let src = sort_token_idx[p] as usize; - for h in 0..hidden { - expected[p * hidden + h] = tokens[src * hidden + h]; - } - } - TestSetup::new(mt_moe_permute::kernel_ir_for(dt)) - .mode(KernelMode::Reduction) - .input(TestBuffer::from_vec("tokens", pack_f32(&tokens_f, dt), dt)) - .input(TestBuffer::from_vec("sort_token_idx", u32_bytes(&sort_token_idx), DType::U32)) - .input(TestBuffer::zeros("permuted", n_permuted * hidden, dt)) - .constexpr("hidden", hidden as u32) - .expect(TestBuffer::from_vec("permuted", pack_f32(&expected, dt), dt)) - .grid_3d(n_permuted as u32, 1, 1, [128, 1, 1]) - } - - // mt_moe_unpermute: weighted sum of k expert outputs back to token order. - #[test_kernel(dtypes = [f32, f16, bf16], tol = [1e-4, 1e-3, 1e-2])] - fn test_moe_unpermute(dt: DType) -> TestSetup { - let (n_tokens, k, hidden) = (4usize, 2usize, 128usize); - let n_permuted = k * n_tokens; - let expert_outputs_f: Vec = - (0..n_permuted * hidden).map(|i| ((i as f32) * 0.011).sin()).collect(); - let inv_perm: Vec = vec![0, 5, 2, 7, 4, 1, 6, 3]; - let weights_f: Vec = (0..n_tokens * k).map(|i| 0.3 + 0.1 * (i as f32)).collect(); - let eo = unpack_f32(&pack_f32(&expert_outputs_f, dt), dt); - let w = unpack_f32(&pack_f32(&weights_f, dt), dt); - let mut expected = vec![0.0f32; n_tokens * hidden]; - for token in 0..n_tokens { - for h in 0..hidden { - let mut acc = 0.0f32; - for j in 0..k { - let pos = inv_perm[token * k + j] as usize; - acc += w[token * k + j] * eo[pos * hidden + h]; - } - expected[token * hidden + h] = acc; - } - } - TestSetup::new(mt_moe_unpermute::kernel_ir_for(dt)) - .mode(KernelMode::Reduction) - .input(TestBuffer::from_vec("expert_outputs", pack_f32(&expert_outputs_f, dt), dt)) - .input(TestBuffer::from_vec("inv_perm", u32_bytes(&inv_perm), DType::U32)) - .input(TestBuffer::from_vec("top_k_weights", pack_f32(&weights_f, dt), dt)) - .input(TestBuffer::zeros("out", n_tokens * hidden, dt)) - .constexpr("hidden", hidden as u32) - .constexpr("k", k as u32) - .expect(TestBuffer::from_vec("out", pack_f32(&expected, dt), dt)) - .grid_3d(n_tokens as u32, 1, 1, [128, 1, 1]) - } } -/// New-syntax benchmarks for the full MoE kernel family. Production-ish -/// Qwen3.6-A3B-ish shapes. All Reduction mode; grids mirror each kernel's -/// DISPATCH INVARIANTS (group counts, never total threads). pub mod kernel_benches { use metaltile::{bench, core::ir::Kernel, test::*}; @@ -4232,75 +3810,4 @@ pub mod kernel_benches { ) } - // ── router_topk — data-dependent argmax, bench-only ─────────────────── - // ABI: router_logits, indices_out, weights_out + {n_experts, k, - // norm_topk_prob}. Grid [B*T, 1, 1], tpg [32,1,1] (pinned in the doc). - #[bench(dtypes = [f32, f16, bf16])] - fn bench_moe_router_topk(dt: DType) -> BenchSetup { - let n_rows = 4096usize; // B*T - let n_experts = 128usize; - let k = 8usize; - let sz = dt.size_bytes(); - let bytes = n_rows * n_experts * sz + n_rows * k * 4 + n_rows * k * sz; - BenchSetup::new(mt_moe_router_topk::kernel_ir_for(dt)) - .mode(KernelMode::Reduction) - .buffer(BenchBuffer::random("router_logits", n_rows * n_experts, dt)) - .buffer(BenchBuffer::zeros("indices_out", n_rows * k, DType::U32).output()) - .buffer(BenchBuffer::zeros("weights_out", n_rows * k, dt).output()) - .constexpr("n_experts", n_experts as u32) - .constexpr("k", k as u32) - .constexpr("norm_topk_prob", 1u32) - .with_shape_label(format!( - "BT{n_rows} E{n_experts} k{k} {}", - crate::utils::dtype_label(dt) - )) - .grid_3d(n_rows as u32, 1, 1, [32, 1, 1]) - .bytes_moved(bytes as u64) - } - - // ── permute — pure gather, bench-only ───────────────────────────────── - // ABI: tokens, sort_token_idx, permuted + {hidden}. Grid [k*B*T, 1, 1], - // tpg [128,1,1]. - #[bench(dtypes = [f32, f16, bf16])] - fn bench_moe_permute(dt: DType) -> BenchSetup { - let bt = 512usize; - let k = 8usize; - let hidden = 2048usize; - let rows = k * bt; - let sz = dt.size_bytes(); - // Reads `rows` source rows (worst case all distinct) + writes `rows`. - let bytes = 2 * rows * hidden * sz; - BenchSetup::new(mt_moe_permute::kernel_ir_for(dt)) - .mode(KernelMode::Reduction) - .buffer(BenchBuffer::random("tokens", bt * hidden, dt)) - .buffer(BenchBuffer::zeros("sort_token_idx", rows, DType::U32)) - .buffer(BenchBuffer::zeros("permuted", rows * hidden, dt).output()) - .constexpr("hidden", hidden as u32) - .with_shape_label(format!("rows{rows} h{hidden} {}", crate::utils::dtype_label(dt))) - .grid_3d(rows as u32, 1, 1, [128, 1, 1]) - .bytes_moved(bytes as u64) - } - - // ── unpermute — weighted scatter-combine, bench-only ────────────────── - // ABI: expert_outputs, inv_perm, top_k_weights, out + {hidden, k}. - // Grid [B*T, 1, 1], tpg [128,1,1]. - #[bench(dtypes = [f32, f16, bf16])] - fn bench_moe_unpermute(dt: DType) -> BenchSetup { - let bt = 512usize; - let k = 8usize; - let hidden = 2048usize; - let sz = dt.size_bytes(); - let bytes = k * bt * hidden * sz + bt * k * 4 + bt * k * sz + bt * hidden * sz; - BenchSetup::new(mt_moe_unpermute::kernel_ir_for(dt)) - .mode(KernelMode::Reduction) - .buffer(BenchBuffer::random("expert_outputs", k * bt * hidden, dt)) - .buffer(BenchBuffer::zeros("inv_perm", bt * k, DType::U32)) - .buffer(BenchBuffer::random("top_k_weights", bt * k, dt)) - .buffer(BenchBuffer::zeros("out", bt * hidden, dt).output()) - .constexpr("hidden", hidden as u32) - .constexpr("k", k as u32) - .with_shape_label(format!("BT{bt} h{hidden} k{k} {}", crate::utils::dtype_label(dt))) - .grid_3d(bt as u32, 1, 1, [128, 1, 1]) - .bytes_moved(bytes as u64) - } } diff --git a/crates/metaltile-std/src/kernels/moe/mod.rs b/crates/metaltile-std/src/kernels/moe/mod.rs index 451c8d43..85f2bf96 100644 --- a/crates/metaltile-std/src/kernels/moe/mod.rs +++ b/crates/metaltile-std/src/kernels/moe/mod.rs @@ -9,11 +9,13 @@ //! //! Filenames drop the redundant `moe_` prefix (the folder provides it); kernel //! names keep `mt_moe_*`. The per-format `*_block_scaled` matrices move as-is; -//! the format-axis fold (plan §7) is deferred. `orchestration.rs` is large and -//! is slated for a follow-up split (router_topk / permute / gather_qmm). +//! the format-axis fold (plan §7) is deferred. The former `orchestration.rs` +//! is split into `router_topk` / `permute` / `gather_qmm`. // Routing — top-k expert selection, permute/unpermute, router pre-scores. -pub mod orchestration; +pub mod router_topk; +pub mod permute; +pub mod gather_qmm; pub mod router_topk_biased; pub mod router_sigmoid_bias; pub mod router_sqrtsoftplus; diff --git a/crates/metaltile-std/src/kernels/moe/permute.rs b/crates/metaltile-std/src/kernels/moe/permute.rs new file mode 100644 index 00000000..3597f445 --- /dev/null +++ b/crates/metaltile-std/src/kernels/moe/permute.rs @@ -0,0 +1,240 @@ +//! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric +//! SPDX-License-Identifier: Apache-2.0 +//! MoE permute / unpermute — `mt_moe_permute` gathers token rows into +//! expert-sorted order ahead of the grouped expert matmul; `mt_moe_unpermute` +//! scatters the per-expert outputs back to token order with the router-weight +//! combine. Both keep the MoE dispatch on-device. + +use metaltile::kernel; + +// ── mt_moe_unpermute ───────────────────────────────────────────────────── +// +// Combine k expert outputs back into the original token order with +// top-k softmax weights. +// +// Inputs: +// expert_outputs — [k*B*T, hidden] per-expert dense outputs at the +// expert-sorted positions +// inv_perm — [B*T, k] where (token i, slot j) was placed +// in expert_outputs (computed by +// caller's sort step) +// top_k_weights — [B*T, k] softmax weights from +// mt_moe_router_topk +// out — [B*T, hidden] weighted sum across k experts +// +// Constexpr: +// hidden — model hidden dim (e.g. 2048 for Qwen3-MoE) +// k — top-k expert count (e.g. 8) +// +// Geometry: +// tpg=128 (split hidden across 128 lanes via 4-wide vectorize) +// grid=[B*T, 1, 1] +// +// Per-token cost: read k * hidden / 128 = (k * hidden) / 128 expert +// values + k weights, do k FMAs per output column, one store per +// column. At hidden=2048, k=8 → ~1k FMAs per token. Bandwidth-bound, +// not ALU-bound. +#[kernel] +pub fn mt_moe_unpermute( + expert_outputs: Tensor, + inv_perm: Tensor, + top_k_weights: Tensor, + mut out: Tensor, + #[constexpr] hidden: u32, + #[constexpr] k: u32, +) { + let token = tgid_x; + let lane = tid; + let row_base_inv = token * k; + let row_base_w = token * k; + let row_base_out = token * hidden; + let n_per_lane = (hidden + 127u32) / 128u32; + for r in range(0u32, n_per_lane, 1u32) { + let h = r * 128u32 + lane; + if h < hidden { + let mut acc = 0.0f32; + for j in range(0u32, k, 1u32) { + let pos = load(inv_perm[row_base_inv + j]); + let v = load(expert_outputs[pos * hidden + h]).cast::(); + let w = load(top_k_weights[row_base_w + j]).cast::(); + acc = acc + w * v; + } + store(out[row_base_out + h], acc.cast::()); + } + } +} + +// ── mt_moe_permute ─────────────────────────────────────────────────────── +// +// Gather tokens into per-expert contiguous buffers given a pre-computed +// sort permutation. The expensive sort step (argsort over top-k expert +// indices) is done by the caller — typically CPU-side via Rust sort, +// or via a future sort kernel. This kernel is just the data-movement +// half: each output position copies the row indicated by sort_token_idx. +// +// Inputs: +// tokens — [B*T, hidden] activations to gather +// sort_token_idx — [k * B*T] for each permuted position p, +// which original token row sourced it. +// Caller computes via argsort over +// top-k indices flattened to +// (token * k + slot) → token (this is +// the "permute" direction; the inverse +// is `inv_perm` consumed by unpermute). +// permuted — [k * B*T, hidden] expert-sorted output. Each k*B*T +// row corresponds to one (expert, token) +// pair; consecutive rows with the same +// expert form that expert's input slab. +// +// Constexpr: +// hidden — model hidden dim +// +// Geometry: +// tpg=128 (split hidden across 128 lanes, ceil(hidden/128) iters) +// grid=[k*B*T, 1, 1] +// +// Per-permuted-row cost: hidden / 128 = 16 loads + 16 stores (at +// hidden=2048). Bandwidth-bound — no FMAs, just a vector copy. +#[kernel] +pub fn mt_moe_permute( + tokens: Tensor, + sort_token_idx: Tensor, + mut permuted: Tensor, + #[constexpr] hidden: u32, +) { + let permuted_pos = tgid_x; + let lane = tid; + let token = load(sort_token_idx[permuted_pos]); + let src_base = token * hidden; + let dst_base = permuted_pos * hidden; + let n_per_lane = (hidden + 127u32) / 128u32; + for r in range(0u32, n_per_lane, 1u32) { + let h = r * 128u32 + lane; + if h < hidden { + let v = load(tokens[src_base + h]); + store(permuted[dst_base + h], v); + } + } +} + + +pub mod kernel_tests { + use metaltile::{test::*, test_kernel}; + + use super::*; + use crate::utils::{pack_f32, unpack_f32}; + + fn u32_bytes(v: &[u32]) -> Vec { v.iter().flat_map(|x| x.to_le_bytes()).collect() } + + // mt_moe_permute: gather token rows into expert-sorted order. + #[test_kernel(dtypes = [f32, f16, bf16], tol = [1e-4, 1e-3, 1e-2])] + fn test_moe_permute(dt: DType) -> TestSetup { + let (n_tokens, k, hidden) = (4usize, 2usize, 128usize); + let n_permuted = k * n_tokens; + let tokens_f: Vec = + (0..n_tokens * hidden).map(|i| ((i as f32) * 0.013).sin()).collect(); + let sort_token_idx: Vec = vec![0, 2, 1, 3, 1, 0, 3, 2]; + let tokens = unpack_f32(&pack_f32(&tokens_f, dt), dt); + let mut expected = vec![0.0f32; n_permuted * hidden]; + for p in 0..n_permuted { + let src = sort_token_idx[p] as usize; + for h in 0..hidden { + expected[p * hidden + h] = tokens[src * hidden + h]; + } + } + TestSetup::new(mt_moe_permute::kernel_ir_for(dt)) + .mode(KernelMode::Reduction) + .input(TestBuffer::from_vec("tokens", pack_f32(&tokens_f, dt), dt)) + .input(TestBuffer::from_vec("sort_token_idx", u32_bytes(&sort_token_idx), DType::U32)) + .input(TestBuffer::zeros("permuted", n_permuted * hidden, dt)) + .constexpr("hidden", hidden as u32) + .expect(TestBuffer::from_vec("permuted", pack_f32(&expected, dt), dt)) + .grid_3d(n_permuted as u32, 1, 1, [128, 1, 1]) + } + + // mt_moe_unpermute: weighted sum of k expert outputs back to token order. + #[test_kernel(dtypes = [f32, f16, bf16], tol = [1e-4, 1e-3, 1e-2])] + fn test_moe_unpermute(dt: DType) -> TestSetup { + let (n_tokens, k, hidden) = (4usize, 2usize, 128usize); + let n_permuted = k * n_tokens; + let expert_outputs_f: Vec = + (0..n_permuted * hidden).map(|i| ((i as f32) * 0.011).sin()).collect(); + let inv_perm: Vec = vec![0, 5, 2, 7, 4, 1, 6, 3]; + let weights_f: Vec = (0..n_tokens * k).map(|i| 0.3 + 0.1 * (i as f32)).collect(); + let eo = unpack_f32(&pack_f32(&expert_outputs_f, dt), dt); + let w = unpack_f32(&pack_f32(&weights_f, dt), dt); + let mut expected = vec![0.0f32; n_tokens * hidden]; + for token in 0..n_tokens { + for h in 0..hidden { + let mut acc = 0.0f32; + for j in 0..k { + let pos = inv_perm[token * k + j] as usize; + acc += w[token * k + j] * eo[pos * hidden + h]; + } + expected[token * hidden + h] = acc; + } + } + TestSetup::new(mt_moe_unpermute::kernel_ir_for(dt)) + .mode(KernelMode::Reduction) + .input(TestBuffer::from_vec("expert_outputs", pack_f32(&expert_outputs_f, dt), dt)) + .input(TestBuffer::from_vec("inv_perm", u32_bytes(&inv_perm), DType::U32)) + .input(TestBuffer::from_vec("top_k_weights", pack_f32(&weights_f, dt), dt)) + .input(TestBuffer::zeros("out", n_tokens * hidden, dt)) + .constexpr("hidden", hidden as u32) + .constexpr("k", k as u32) + .expect(TestBuffer::from_vec("out", pack_f32(&expected, dt), dt)) + .grid_3d(n_tokens as u32, 1, 1, [128, 1, 1]) + } +} + +pub mod kernel_benches { + use metaltile::{bench, test::*}; + + use super::*; + + // ── permute — pure gather, bench-only ───────────────────────────────── + // ABI: tokens, sort_token_idx, permuted + {hidden}. Grid [k*B*T, 1, 1], + // tpg [128,1,1]. + #[bench(dtypes = [f32, f16, bf16])] + fn bench_moe_permute(dt: DType) -> BenchSetup { + let bt = 512usize; + let k = 8usize; + let hidden = 2048usize; + let rows = k * bt; + let sz = dt.size_bytes(); + // Reads `rows` source rows (worst case all distinct) + writes `rows`. + let bytes = 2 * rows * hidden * sz; + BenchSetup::new(mt_moe_permute::kernel_ir_for(dt)) + .mode(KernelMode::Reduction) + .buffer(BenchBuffer::random("tokens", bt * hidden, dt)) + .buffer(BenchBuffer::zeros("sort_token_idx", rows, DType::U32)) + .buffer(BenchBuffer::zeros("permuted", rows * hidden, dt).output()) + .constexpr("hidden", hidden as u32) + .with_shape_label(format!("rows{rows} h{hidden} {}", crate::utils::dtype_label(dt))) + .grid_3d(rows as u32, 1, 1, [128, 1, 1]) + .bytes_moved(bytes as u64) + } + + // ── unpermute — weighted scatter-combine, bench-only ────────────────── + // ABI: expert_outputs, inv_perm, top_k_weights, out + {hidden, k}. + // Grid [B*T, 1, 1], tpg [128,1,1]. + #[bench(dtypes = [f32, f16, bf16])] + fn bench_moe_unpermute(dt: DType) -> BenchSetup { + let bt = 512usize; + let k = 8usize; + let hidden = 2048usize; + let sz = dt.size_bytes(); + let bytes = k * bt * hidden * sz + bt * k * 4 + bt * k * sz + bt * hidden * sz; + BenchSetup::new(mt_moe_unpermute::kernel_ir_for(dt)) + .mode(KernelMode::Reduction) + .buffer(BenchBuffer::random("expert_outputs", k * bt * hidden, dt)) + .buffer(BenchBuffer::zeros("inv_perm", bt * k, DType::U32)) + .buffer(BenchBuffer::random("top_k_weights", bt * k, dt)) + .buffer(BenchBuffer::zeros("out", bt * hidden, dt).output()) + .constexpr("hidden", hidden as u32) + .constexpr("k", k as u32) + .with_shape_label(format!("BT{bt} h{hidden} k{k} {}", crate::utils::dtype_label(dt))) + .grid_3d(bt as u32, 1, 1, [128, 1, 1]) + .bytes_moved(bytes as u64) + } +} diff --git a/crates/metaltile-std/src/kernels/moe/router_topk.rs b/crates/metaltile-std/src/kernels/moe/router_topk.rs new file mode 100644 index 00000000..d0725252 --- /dev/null +++ b/crates/metaltile-std/src/kernels/moe/router_topk.rs @@ -0,0 +1,274 @@ +//! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric +//! SPDX-License-Identifier: Apache-2.0 +//! MoE router top-k expert selection — `mt_moe_router_topk` picks the top-k +//! experts per token by logit and emits the normalised routing weights +//! (softmax over chosen-k, or global softmax), on-device so routing never +//! round-trips to the CPU. The biased/unbiased dual-score variant lives in +//! `router_topk_biased.rs`. + +use metaltile::kernel; + +// ── mt_moe_router_topk ─────────────────────────────────────────────────── +// +// Per-token select top-k experts from `router_logits`, plus softmax +// weights over the chosen k. +// +// Inputs: +// router_logits — [B*T, n_experts] (any float dtype, computed in f32) +// indices_out — [B*T, k] (u32) +// weights_out — [B*T, k] (same dtype as router_logits, softmax weights) +// +// Constexpr: +// n_experts — typical Qwen3.6-A3B: 128. Must fit one simdgroup +// (≤ 32×32 = 1024) — every reasonable MoE topology. +// k — typical 6-8 for production MoE. Hard cap k ≤ 32. +// +// Geometry: +// tpg=32 (one simdgroup per token row) +// grid = [B*T, 1, 1] (Reduction mode) +// +// Algorithm — k iterations of simd-parallel argmax with mask of +// previously-chosen indices stored in TG memory. After k passes, +// softmax over the chosen k values in-place on lane 0..k-1. +// +// Bench spec uses BenchDispatch::Generic + shapes: &[] so `tile bench` +// skips it; correctness lives in unit tests + downstream MoE +// integration. Same convention as other ffai/ kernels (gather, sampling). +#[kernel] +pub fn mt_moe_router_topk( + router_logits: Tensor, + mut indices_out: Tensor, + mut weights_out: Tensor, + #[constexpr] n_experts: u32, + #[constexpr] k: u32, + // 1 = Qwen3-MoE style (softmax over chosen-k, sum-to-1 — `norm_topk_prob=True`) + // 0 = Qwen3-Next style (softmax over ALL n_experts, return chosen probs + // un-renormalized — `norm_topk_prob=False`) + // Mathematically equivalent at mode 1: softmax-over-chosen-k is the + // same as (softmax-over-all → renormalize-over-chosen). Mode 0 + // returns probs that sum to < 1 across the chosen k, matching MLX's + // qwen3_next.py:334-341. + // + // INVARIANT: this kernel pins tpg=32 (one simdgroup per token row). + // The `simdgroup_barrier_mem_none()` below is correct only at tpg=32. + // Caller must dispatch with `[n_rows, 1, 1] × [32, 1, 1]`. + #[constexpr] norm_topk_prob: u32, +) { + let row = tgid_x; + let lane = tid; + let row_base = row * n_experts; + // TG scratch: chosen indices + values from each of the k argmax passes. + // 32 slots covers any reasonable k (typical 6-8). Kernel assumes + // k ≤ 32 — caller MUST enforce this in the host-side dispatcher + // (no GPU-side check, would silently scribble into adjacent TG mem). + threadgroup_alloc("tg_chosen_idx", 32u32); + threadgroup_alloc("tg_chosen_val", 32u32); + // Cache the all-experts-softmax sum for Qwen3-Next mode (mode 0). + // 1 slot, written by lane 0 in the prepass. + threadgroup_alloc("tg_full_sum", 1u32); + threadgroup_alloc("tg_full_max", 1u32); + // ── Pre-pass: compute softmax denominator over ALL n_experts ───── + // Needed only for norm_topk_prob=0 (Qwen3-Next), but the cost is + // trivial (one simd_max + simd_sum) and emitting it unconditionally + // keeps the codegen tight (the codegen DCE will drop the dead path + // when the constexpr branch is unreachable). + let mut local_max_all = neg_infinity(); + let n_per_lane_pre = (n_experts + 31u32) / 32u32; + for r in range(0u32, n_per_lane_pre, 1u32) { + let j = r * 32u32 + lane; + if j < n_experts { + let v = load(router_logits[row_base + j]).cast::(); + let better = v > local_max_all; + local_max_all = select(better, v, local_max_all); + } + } + let row_max_all = simd_max(local_max_all); + let mut local_sum_all = 0.0f32; + for r in range(0u32, n_per_lane_pre, 1u32) { + let j = r * 32u32 + lane; + if j < n_experts { + let v = load(router_logits[row_base + j]).cast::(); + local_sum_all = local_sum_all + exp(v - row_max_all); + } + } + let row_sum_all = simd_sum(local_sum_all); + if lane == 0u32 { + threadgroup_store("tg_full_max", 0u32, row_max_all); + threadgroup_store("tg_full_sum", 0u32, row_sum_all); + } + simdgroup_barrier_mem_none(); + // ── k argmax passes with chosen-mask ───────────────────────────── + for it in range(0u32, k, 1u32) { + // Per-lane local argmax over its slice of n_experts. + // Each lane covers ceil(n_experts/32) experts. + let mut best_val = neg_infinity(); + let mut best_idx = 0u32; + let n_per_lane = (n_experts + 31u32) / 32u32; + for r in range(0u32, n_per_lane, 1u32) { + let j = r * 32u32 + lane; + if j < n_experts { + let v = load(router_logits[row_base + j]).cast::(); + // Mask: was j picked in a previous iter? + // Scan tg_chosen_idx[0..it] — k ≤ 8 typically so this + // is fast even without early exit. + let mut chosen_mask = 0u32; + for p in range(0u32, it, 1u32) { + let cp = threadgroup_load("tg_chosen_idx", p); + chosen_mask = chosen_mask | select(j == cp, 1u32, 0u32); + } + let candidate = select(chosen_mask > 0u32, neg_infinity(), v); + let better = candidate > best_val; + best_val = select(better, candidate, best_val); + best_idx = select(better, j, best_idx); + } + } + // Cross-lane reduce. simd_max gives the global best value; + // ties broken to smaller idx via simd_min on (idx | sentinel). + let global_best_val = simd_max(best_val); + let i_have = best_val == global_best_val; + let my_idx_or_max = select(i_have, best_idx, 4294967295u32); // u32::MAX + let global_best_idx = simd_min(my_idx_or_max); + // Lane 0 writes the iter's chosen slot. + if lane == 0u32 { + threadgroup_store("tg_chosen_idx", it, global_best_idx); + threadgroup_store("tg_chosen_val", it, global_best_val); + } + simdgroup_barrier_mem_none(); + } + // ── Softmax / weight emit per `norm_topk_prob` ────────────────── + // Mode 1 (Qwen3-MoE, default): softmax over chosen-k (sum-to-1). + // numerator = exp(z_i - max_chosen); divisor = Σ_j∈chosen + // == exp(z_i - max_all) · const / Σ_j∈chosen exp(z_j - max_all) · const + // so we can use the SAME numerator as mode 0 (exp(z - max_all)) and + // just swap the divisor. Avoids needing a Rust `if`-expression + // which the DSL doesn't unify across arms. + // Mode 0 (Qwen3-Next): un-normalized chosen probs (sum < 1). + // weight_i = exp(z_i - max_all) / Σ_j∈all exp(z_j - max_all) + let my_val = select(lane < k, threadgroup_load("tg_chosen_val", lane), neg_infinity()); + let row_max_full = threadgroup_load("tg_full_max", 0u32); + let row_sum_full = threadgroup_load("tg_full_sum", 0u32); + let exp_val = exp(my_val - row_max_full); + let masked_exp = select(lane < k, exp_val, 0.0f32); + let sum_chosen = simd_sum(masked_exp); + // Pick divisor: chosen-k sum for renormalized (mode 1) or all-experts + // sum for raw probs (mode 0). select() forces both to be live; codegen + // const-folds when `norm_topk_prob` bakes in. + let divisor = select(norm_topk_prob == 1u32, sum_chosen, row_sum_full); + let weight = masked_exp / divisor; + // ── Write outputs ─────────────────────────────────────────────── + if lane < k { + let out_base = row * k + lane; + store(indices_out[out_base], threadgroup_load("tg_chosen_idx", lane)); + store(weights_out[out_base], weight.cast::()); + } +} + + +pub mod kernel_tests { + use metaltile::{test::*, test_kernel}; + + use super::*; + use crate::utils::{pack_f32, unpack_f32}; + + fn u32_bytes(v: &[u32]) -> Vec { v.iter().flat_map(|x| x.to_le_bytes()).collect() } + + // ── Router top-k selection ───────────────────────────────────────────────────────────── + + /// Router oracle: softmax over all experts for the denominator, pick the + /// top-k logits (well-separated test inputs → no ties), then weight either + /// by renormalised softmax over the chosen k (`norm_topk_prob`) or by the + /// global softmax (raw probs that sum to < 1). + fn router_oracle( + logits: &[f32], + n_rows: usize, + n_experts: usize, + k: usize, + norm_topk_prob: bool, + ) -> (Vec, Vec) { + let mut idx_out = vec![0u32; n_rows * k]; + let mut w_out = vec![0.0f32; n_rows * k]; + for row in 0..n_rows { + let row_l = &logits[row * n_experts..(row + 1) * n_experts]; + let max_all = row_l.iter().copied().fold(f32::NEG_INFINITY, f32::max); + let sum_all: f32 = row_l.iter().map(|&l| (l - max_all).exp()).sum(); + // Top-k by descending logit (stable: smaller index wins ties). + let mut order: Vec = (0..n_experts).collect(); + order.sort_by(|&a, &b| row_l[b].partial_cmp(&row_l[a]).unwrap().then(a.cmp(&b))); + let chosen = &order[..k]; + let sum_chosen: f32 = chosen.iter().map(|&e| (row_l[e] - max_all).exp()).sum(); + for (i, &e) in chosen.iter().enumerate() { + idx_out[row * k + i] = e as u32; + let num = (row_l[e] - max_all).exp(); + w_out[row * k + i] = if norm_topk_prob { num / sum_chosen } else { num / sum_all }; + } + } + (idx_out, w_out) + } + + fn router_setup(dt: DType, norm_topk_prob: bool) -> TestSetup { + let (n_rows, n_experts, k) = (4usize, 8usize, 4usize); + // Well-separated logits (distinct multiples of 0.5 per row → no ties, + // gap ≫ dtype epsilon so the selection is dtype-stable). + let logits_f: Vec = (0..n_rows * n_experts) + .map(|i| { + let row = i / n_experts; + let e = i % n_experts; + ((e * 5 + row * 3) % n_experts) as f32 * 0.5 + }) + .collect(); + let logits = unpack_f32(&pack_f32(&logits_f, dt), dt); + let (idx, w) = router_oracle(&logits, n_rows, n_experts, k, norm_topk_prob); + TestSetup::new(mt_moe_router_topk::kernel_ir_for(dt)) + .mode(KernelMode::Reduction) + .input(TestBuffer::from_vec("router_logits", pack_f32(&logits_f, dt), dt)) + .input(TestBuffer::zeros("indices_out", n_rows * k, DType::U32)) + .input(TestBuffer::zeros("weights_out", n_rows * k, dt)) + .constexpr("n_experts", n_experts as u32) + .constexpr("k", k as u32) + .constexpr("norm_topk_prob", u32::from(norm_topk_prob)) + .expect(TestBuffer::from_vec("indices_out", u32_bytes(&idx), DType::U32)) + .expect(TestBuffer::from_vec("weights_out", pack_f32(&w, dt), dt)) + .grid_3d(n_rows as u32, 1, 1, [32, 1, 1]) + } + + // norm_topk_prob = 1: weights renormalised over the chosen k (Qwen3-MoE). + #[test_kernel(dtypes = [f32, f16, bf16], tol = [1e-4, 1e-2, 5e-2])] + fn test_moe_router_topk_norm(dt: DType) -> TestSetup { router_setup(dt, true) } + // norm_topk_prob = 0: raw global-softmax probs (Qwen3-Next). + #[test_kernel(dtypes = [f32, f16, bf16], tol = [1e-4, 1e-2, 5e-2])] + fn test_moe_router_topk_global(dt: DType) -> TestSetup { router_setup(dt, false) } + +} + +pub mod kernel_benches { + use metaltile::{bench, test::*}; + + use super::*; + + // ── router_topk — data-dependent argmax, bench-only ─────────────────── + // ABI: router_logits, indices_out, weights_out + {n_experts, k, + // norm_topk_prob}. Grid [B*T, 1, 1], tpg [32,1,1] (pinned in the doc). + #[bench(dtypes = [f32, f16, bf16])] + fn bench_moe_router_topk(dt: DType) -> BenchSetup { + let n_rows = 4096usize; // B*T + let n_experts = 128usize; + let k = 8usize; + let sz = dt.size_bytes(); + let bytes = n_rows * n_experts * sz + n_rows * k * 4 + n_rows * k * sz; + BenchSetup::new(mt_moe_router_topk::kernel_ir_for(dt)) + .mode(KernelMode::Reduction) + .buffer(BenchBuffer::random("router_logits", n_rows * n_experts, dt)) + .buffer(BenchBuffer::zeros("indices_out", n_rows * k, DType::U32).output()) + .buffer(BenchBuffer::zeros("weights_out", n_rows * k, dt).output()) + .constexpr("n_experts", n_experts as u32) + .constexpr("k", k as u32) + .constexpr("norm_topk_prob", 1u32) + .with_shape_label(format!( + "BT{n_rows} E{n_experts} k{k} {}", + crate::utils::dtype_label(dt) + )) + .grid_3d(n_rows as u32, 1, 1, [32, 1, 1]) + .bytes_moved(bytes as u64) + } + +} diff --git a/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs index 9cb0d3fc..2b82d63b 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs @@ -24,7 +24,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::orchestration::{ +use metaltile_std::kernels::moe::gather_qmm::{ mt_moe_gather_qmm_int4_m8, mt_moe_gather_qmm_int4_m16, mt_moe_gather_qmm_int4_m32, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs b/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs index 61deab89..4625ff0d 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs @@ -13,7 +13,7 @@ use std::{collections::BTreeMap, time::Instant}; use common::{Dt, gpu_lock, pack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::orchestration::{ +use metaltile_std::kernels::moe::gather_qmm::{ mt_moe_gather_qmm_int4, mt_moe_gather_qmm_int4_m8, mt_moe_gather_qmm_mma_int4, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs index 0fb80c12..e4056cc4 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs @@ -3,7 +3,7 @@ #![allow(clippy::manual_is_multiple_of)] //! GPU correctness for the bit-width-generalized MMA MoE BGEMMs -//! `kernels::moe::orchestration::mt_moe_gather_qmm_mma_b{3,5,6,8}`. +//! `kernels::moe::gather_qmm::mt_moe_gather_qmm_mma_b{3,5,6,8}`. //! //! Same tiled-MMA algorithm as `mt_moe_gather_qmm_mma_int4`, but the //! weight coop-dequant pulls codes from a contiguous LSB-first bit-stream @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::orchestration::{ +use metaltile_std::kernels::moe::gather_qmm::{ mt_moe_gather_qmm_mma_b3, mt_moe_gather_qmm_mma_b5, mt_moe_gather_qmm_mma_b6, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs index f755ccdc..e1a82ad3 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_int4, mpp_bm64}; +use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_int4, mpp_bm64}; /// Pack a row of int4 weights into uint32s (8 per uint, LSB-first per /// nibble). Identical to the helper used by the bm16_mpp test — diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs index 997b433d..4efc3bad 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs @@ -21,7 +21,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_b8, mpp_bm64_int8}; +use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_b8, mpp_bm64_int8}; /// Pack a row of int8 weight codes into uint32s (4 codes per uint, LE byte /// order). Code values must be in 0..=255. diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs index 163cc28a..bbff0b09 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs @@ -23,7 +23,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_int4, mpp_bm8}; +use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_int4, mpp_bm8}; /// Pack a row of int4 weights into uint32s (8 per uint, LSB-first per /// nibble). Same helper used by the bm16_mpp / bm64_mpp test files — diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs index 8f6e99b2..a6be61be 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_b8, mpp_bm8_int8}; +use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_b8, mpp_bm8_int8}; /// Pack a row of int8 weight codes into uint32s (4 codes per uint, LE byte /// order). Code values must be in 0..=255. diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs index f2ea2d26..e4e8f18f 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs @@ -23,7 +23,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{orchestration::mt_moe_gather_qmm_int4, mpp}; +use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_int4, mpp}; /// Pack a row of int4 weights into uint32s (8 per uint, LSB-first per nibble). /// Identical to the helper used by the legacy diff --git a/docs/specs/KERNEL_AUDIT.md b/docs/specs/KERNEL_AUDIT.md index e8d349d8..fc1f9677 100644 --- a/docs/specs/KERNEL_AUDIT.md +++ b/docs/specs/KERNEL_AUDIT.md @@ -104,8 +104,8 @@ Local verification of NAX kernels is the developer's responsibility on M4+ hardw | gemv_masked | ✓ | ✓ | ✓ | `kernels/gemm/gemv_masked.rs` → `mt_gemv_masked`. | | quantized (affine_quantize / affine_dequantize) | ✓ | ✓ | ✓ | `kernels/gemm/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) | ✓ | ✓ | ✓ | `kernels/gemm/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) | ✓ | ✓ | ✓ | `kernels/moe/orchestration.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` (`kernels/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): `mt_dequant_gemv_int4_expert_indexed` — per-output-row expert selection for the gate/up FFN dispatch shape. | -| moe (router top-k + permute + unpermute orchestration) | ✗ | ✓ | ✓ | `kernels/moe/orchestration.rs` → `mt_moe_router_topk`, `mt_moe_permute`, `mt_moe_unpermute`. MoE expert-routing orchestration. The grouped quantized BGEMM that fuses per-expert FFN matmuls is counted under the `quantized (gather_*)` row. | +| quantized (gather_qmv / gather_qmm — gather variants) | ✓ | ✓ | ✓ | `kernels/moe/gather_qmm.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` (`kernels/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): `mt_dequant_gemv_int4_expert_indexed` — per-output-row expert selection for the gate/up FFN dispatch shape. | +| moe (router top-k + permute + unpermute orchestration) | ✗ | ✓ | ✓ | `kernels/moe/router_topk.rs` (`mt_moe_router_topk`) + `kernels/moe/permute.rs` (`mt_moe_permute`, `mt_moe_unpermute`). MoE expert-routing orchestration. The grouped quantized BGEMM that fuses per-expert FFN matmuls is counted under the `quantized (gather_*)` row. | | dequant_gather (quantized embedding-table gather) | ✗ | ✗ | ✓ | `ffai/dequant_gather.rs`. int{3,4,5,6,8} all bit-widths. FFAI-only. | | dequant_gemv (quantized GEMV, FFAI flavour) | ~ | ~ | ✓ | `kernels/gemm/dequant_gemv.rs` → `mt_dequant_gemv_int{2,3,4,5,6,8}` (one-row-per-TG) + `mt_dequant_gemv_int4_fast` (PR #154, 8-row-per-TG, mirrors MLX `qmv_fast`). The non-fast int4 kernel stays because FFAI's GPU-router opts into its indirect Swift wrapper. | | fp_quantized (fp4/fp8 quant + dequant) | ✓ | ✓ | ✓ | `kernels/gemm/fp_quantized.rs` → `mt_fp4_quant_dequant` (fp4 E2M1) + `mt_fp8_e4m3_quant_dequant` / `mt_fp8_e5m2_quant_dequant` (fp8). Pure arithmetic transform (per-group max-scale + mantissa rounding via `floor(log2)`/`exp2`/`round`); exact for fp8 normals/subnormals, saturating (no NaN/Inf). | @@ -335,7 +335,7 @@ A few rows mix multiple `.metal` files into one op or split one file into multip - **`steel/`** — each kernel file becomes one op row; per-block-shape instantiations are not counted separately. `steel_attention` (scalar) and `steel_attention_mma` (simdgroup-MMA) are two rows because they are separately compiled kernels with different lowering strategies. - **`quantized.metal`** — split into four rows by semantic operation (quant/dequant, qmv/qvm/qmm matmul, gather-qmv/qmm, fp4/fp8). The Apple10+ variants (`quantized_nax`, `fp_quantized_nax`) are separate rows because they live in separate modules with runtime-only dispatch gating. `fp_quantized_mma` is its own row (runs on M1+, no Apple10 gating). - **`indexing/`** is one row covering scatter / scatter_axis / gather_axis / gather_front / masked_scatter. Bare `gather` is its own row (FFAI-specific). -- **`moe`** is the routing/permute/unpermute orchestration in `kernels/moe/orchestration.rs`. The grouped quantized BGEMM lives under the `quantized (gather_*)` row. +- **`moe`** is the routing (`kernels/moe/router_topk.rs`) + permute/unpermute (`kernels/moe/permute.rs`). The grouped quantized BGEMM lives under the `quantized (gather_*)` row. - **`logits processors`** is one row for the FFAI sampler-stage kernels (`temperature`, `repetition_penalty`, `topk` / `top_p` / `min_p` masks). - Cells marked **`~`** indicate a partial port (typically one bit-width, one dtype, or one block shape where upstream has many) — see the notes column for the specific gap. diff --git a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md index 9a64f37e..57928528 100644 --- a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md +++ b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md @@ -55,12 +55,12 @@ crates/metaltile-std/src/kernels/ 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 - moe/ ✅ DONE — orchestration (router_topk + permute/unpermute + gather_qmm) · + moe/ ✅ DONE — router_topk · permute (+unpermute) · gather_qmm (per-expert BGEMM) · router_topk_biased / sigmoid_bias / sqrtsoftplus · mpp(bm8/bm64 × int8 × block_scaled) + mpp_shared · bgemm/gemv (q2k/iq2xxs/q4, view/ws/rows) · gather_q4 · down_swiglu_accum / down_weighted_sum · dequant_gemv_expert_indexed(_block_scaled) · block_scaled. Filenames drop the moe_ prefix; format-axis fold deferred (§7). - orchestration.rs (~4k lines) slated for a follow-up split. + (orchestration split into router_topk / permute / gather_qmm.) norm/ ✅ DONE — rms_norm(+residual/rope/qgemv/gated) · layer_norm · adain1d rope/ ✅ DONE — rope · rope_2d · rope_banded · rope_yarn · partial_rope convolution/ ✅ DONE — conv1d/2d/3d · depthwise · winograd · steel_conv · conv1d_causal(_roll) (see §4) From 4156eecb398c04c74e649a2acad6417dbf1d5d49 Mon Sep 17 00:00:00 2001 From: Eric Kryski <599019+ekryski@users.noreply.github.com> Date: Sat, 13 Jun 2026 14:03:42 -0600 Subject: [PATCH 3/7] refactor(moe): restore moe_ filename prefix to match kernel names MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Revert the prefix-drop from the moe migration: filenames carry the moe_ prefix again so they match their mt_moe_* kernel names (moe_mpp.rs, moe_gather_qmm.rs, moe_router_topk.rs, moe_gather_q4.rs, …). The two files whose kernels are format-prefixed instead keep names matching those (block_scaled_moe.rs -> mt__gather_qmm; dequant_gemv_expert_indexed*). Updated mod.rs + all consumer-test imports (grouped + mixed use-blocks). --- .../{block_scaled.rs => block_scaled_moe.rs} | 0 crates/metaltile-std/src/kernels/moe/mod.rs | 85 ++++++++++--------- ...q2xxs_bm64.rs => moe_bgemm_iq2xxs_bm64.rs} | 0 ..._iq2xxs_mpp.rs => moe_bgemm_iq2xxs_mpp.rs} | 0 ...q2xxs_view.rs => moe_bgemm_iq2xxs_view.rs} | 0 ...4.rs => moe_bgemm_iq2xxs_view_u16_bm64.rs} | 0 ...gemm_q2k_bm64.rs => moe_bgemm_q2k_bm64.rs} | 0 ...{bgemm_q2k_mpp.rs => moe_bgemm_q2k_mpp.rs} | 0 ...gemm_q2k_view.rs => moe_bgemm_q2k_view.rs} | 0 ...bm64.rs => moe_bgemm_q2k_view_u16_bm64.rs} | 0 ...{bgemm_q4_bm64.rs => moe_bgemm_q4_bm64.rs} | 0 ...iglu_accum.rs => moe_down_swiglu_accum.rs} | 0 ...um_f16.rs => moe_down_weighted_sum_f16.rs} | 0 ...her_down_q2k.rs => moe_gather_down_q2k.rs} | 0 ...mv_iq2xxs.rs => moe_gather_gemv_iq2xxs.rs} | 0 .../moe/{gather_q4.rs => moe_gather_q4.rs} | 0 .../moe/{gather_qmm.rs => moe_gather_qmm.rs} | 0 ...rows_iq2xxs.rs => moe_gemv_rows_iq2xxs.rs} | 0 ...{gemv_rows_q2k.rs => moe_gemv_rows_q2k.rs} | 0 ...iq2xxs.rs => moe_gemv_rows_view_iq2xxs.rs} | 0 ...emv_ws_iq2xxs.rs => moe_gemv_ws_iq2xxs.rs} | 0 .../{gemv_ws_q2k.rs => moe_gemv_ws_q2k.rs} | 0 .../src/kernels/moe/{mpp.rs => moe_mpp.rs} | 4 +- ...lock_scaled.rs => moe_mpp_block_scaled.rs} | 0 .../moe/{mpp_bm64.rs => moe_mpp_bm64.rs} | 4 +- ...scaled.rs => moe_mpp_bm64_block_scaled.rs} | 0 ...{mpp_bm64_int8.rs => moe_mpp_bm64_int8.rs} | 4 +- .../moe/{mpp_bm8.rs => moe_mpp_bm8.rs} | 4 +- ..._scaled.rs => moe_mpp_bm8_block_scaled.rs} | 0 .../{mpp_bm8_int8.rs => moe_mpp_bm8_int8.rs} | 4 +- .../moe/{mpp_int8.rs => moe_mpp_int8.rs} | 4 +- .../moe/{mpp_shared.rs => moe_mpp_shared.rs} | 0 .../moe/{permute.rs => moe_permute.rs} | 0 ...oid_bias.rs => moe_router_sigmoid_bias.rs} | 0 ...softplus.rs => moe_router_sqrtsoftplus.rs} | 0 .../{router_topk.rs => moe_router_topk.rs} | 0 ...pk_biased.rs => moe_router_topk_biased.rs} | 0 .../{sigmoid_bias.rs => moe_sigmoid_bias.rs} | 0 .../tests/bm64_vs_gemvrows_iq2xxs.rs | 4 +- .../tests/dsv4_router_topk_correctness.rs | 2 +- .../moe_bgemm_iq2xxs_bm64_correctness.rs | 4 +- .../tests/moe_bgemm_iq2xxs_mpp_correctness.rs | 2 +- .../moe_bgemm_iq2xxs_view_correctness.rs | 2 +- .../tests/moe_bgemm_q2k_bm64_correctness.rs | 4 +- .../tests/moe_bgemm_q2k_mpp_correctness.rs | 2 +- .../tests/moe_bgemm_q2k_view_correctness.rs | 4 +- .../tests/moe_bm64_ragged_correctness.rs | 4 +- .../tests/moe_gather_down_q2k_correctness.rs | 4 +- .../moe_gather_gemv_iq2xxs_correctness.rs | 4 +- ...moe_gather_qmm_int4_m16_m32_correctness.rs | 2 +- .../tests/moe_gather_qmm_microbench.rs | 2 +- ...moe_gather_qmm_mma_bitwidth_correctness.rs | 4 +- .../moe_gather_qmm_mpp_bm64_correctness.rs | 10 +-- ...oe_gather_qmm_mpp_bm64_int8_correctness.rs | 10 +-- .../moe_gather_qmm_mpp_bm8_correctness.rs | 6 +- ...moe_gather_qmm_mpp_bm8_int8_correctness.rs | 6 +- .../tests/moe_gather_qmm_mpp_correctness.rs | 8 +- .../moe_gather_qmm_mpp_int8_correctness.rs | 6 +- .../tests/moe_gemv_rows_iq2xxs_correctness.rs | 2 +- .../tests/moe_gemv_rows_q2k_correctness.rs | 4 +- .../moe_gemv_rows_view_iq2xxs_correctness.rs | 2 +- .../tests/moe_gemv_ws_iq2xxs_correctness.rs | 2 +- .../tests/moe_gemv_ws_q2k_correctness.rs | 4 +- .../tests/moe_q2k_view_u16_correctness.rs | 4 +- .../tests/moe_view_u16_correctness.rs | 4 +- docs/specs/KERNEL_CONSOLIDATION_PLAN.md | 2 +- 66 files changed, 112 insertions(+), 111 deletions(-) rename crates/metaltile-std/src/kernels/moe/{block_scaled.rs => block_scaled_moe.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{bgemm_iq2xxs_bm64.rs => moe_bgemm_iq2xxs_bm64.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{bgemm_iq2xxs_mpp.rs => moe_bgemm_iq2xxs_mpp.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{bgemm_iq2xxs_view.rs => moe_bgemm_iq2xxs_view.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{bgemm_iq2xxs_view_u16_bm64.rs => moe_bgemm_iq2xxs_view_u16_bm64.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{bgemm_q2k_bm64.rs => moe_bgemm_q2k_bm64.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{bgemm_q2k_mpp.rs => moe_bgemm_q2k_mpp.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{bgemm_q2k_view.rs => moe_bgemm_q2k_view.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{bgemm_q2k_view_u16_bm64.rs => moe_bgemm_q2k_view_u16_bm64.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{bgemm_q4_bm64.rs => moe_bgemm_q4_bm64.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{down_swiglu_accum.rs => moe_down_swiglu_accum.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{down_weighted_sum_f16.rs => moe_down_weighted_sum_f16.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{gather_down_q2k.rs => moe_gather_down_q2k.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{gather_gemv_iq2xxs.rs => moe_gather_gemv_iq2xxs.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{gather_q4.rs => moe_gather_q4.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{gather_qmm.rs => moe_gather_qmm.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{gemv_rows_iq2xxs.rs => moe_gemv_rows_iq2xxs.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{gemv_rows_q2k.rs => moe_gemv_rows_q2k.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{gemv_rows_view_iq2xxs.rs => moe_gemv_rows_view_iq2xxs.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{gemv_ws_iq2xxs.rs => moe_gemv_ws_iq2xxs.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{gemv_ws_q2k.rs => moe_gemv_ws_q2k.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{mpp.rs => moe_mpp.rs} (98%) rename crates/metaltile-std/src/kernels/moe/{mpp_block_scaled.rs => moe_mpp_block_scaled.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{mpp_bm64.rs => moe_mpp_bm64.rs} (98%) rename crates/metaltile-std/src/kernels/moe/{mpp_bm64_block_scaled.rs => moe_mpp_bm64_block_scaled.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{mpp_bm64_int8.rs => moe_mpp_bm64_int8.rs} (98%) rename crates/metaltile-std/src/kernels/moe/{mpp_bm8.rs => moe_mpp_bm8.rs} (98%) rename crates/metaltile-std/src/kernels/moe/{mpp_bm8_block_scaled.rs => moe_mpp_bm8_block_scaled.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{mpp_bm8_int8.rs => moe_mpp_bm8_int8.rs} (98%) rename crates/metaltile-std/src/kernels/moe/{mpp_int8.rs => moe_mpp_int8.rs} (98%) rename crates/metaltile-std/src/kernels/moe/{mpp_shared.rs => moe_mpp_shared.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{permute.rs => moe_permute.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{router_sigmoid_bias.rs => moe_router_sigmoid_bias.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{router_sqrtsoftplus.rs => moe_router_sqrtsoftplus.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{router_topk.rs => moe_router_topk.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{router_topk_biased.rs => moe_router_topk_biased.rs} (100%) rename crates/metaltile-std/src/kernels/moe/{sigmoid_bias.rs => moe_sigmoid_bias.rs} (100%) diff --git a/crates/metaltile-std/src/kernels/moe/block_scaled.rs b/crates/metaltile-std/src/kernels/moe/block_scaled_moe.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/block_scaled.rs rename to crates/metaltile-std/src/kernels/moe/block_scaled_moe.rs diff --git a/crates/metaltile-std/src/kernels/moe/mod.rs b/crates/metaltile-std/src/kernels/moe/mod.rs index 85f2bf96..0b479ac3 100644 --- a/crates/metaltile-std/src/kernels/moe/mod.rs +++ b/crates/metaltile-std/src/kernels/moe/mod.rs @@ -7,59 +7,60 @@ //! BGEMM/GEMV, batched Q4 gather), and the down-projection combine. Migrated //! from the legacy `mlx/` + `ffai/` split. //! -//! Filenames drop the redundant `moe_` prefix (the folder provides it); kernel -//! names keep `mt_moe_*`. The per-format `*_block_scaled` matrices move as-is; -//! the format-axis fold (plan §7) is deferred. The former `orchestration.rs` -//! is split into `router_topk` / `permute` / `gather_qmm`. +//! Filenames keep the `moe_` prefix so they match their kernel names +//! (`mt_moe_*`); the two files whose kernels are format-prefixed instead +//! (`block_scaled_moe`, `dequant_gemv_expert_indexed*`) match those. The +//! per-format `*_block_scaled` matrices keep the format axis as-is; the +//! format-axis fold (plan §7) is deferred. // Routing — top-k expert selection, permute/unpermute, router pre-scores. -pub mod router_topk; -pub mod permute; -pub mod gather_qmm; -pub mod router_topk_biased; -pub mod router_sigmoid_bias; -pub mod router_sqrtsoftplus; -pub mod sigmoid_bias; +pub mod moe_router_topk; +pub mod moe_permute; +pub mod moe_gather_qmm; +pub mod moe_router_topk_biased; +pub mod moe_router_sigmoid_bias; +pub mod moe_router_sqrtsoftplus; +pub mod moe_sigmoid_bias; // MPP grouped BGEMM (one ABI, tile-geometry / bit-width variants; shared -// test/bench helpers in `mpp_shared`). -pub mod mpp; -pub mod mpp_shared; -pub mod mpp_int8; -pub mod mpp_bm8; -pub mod mpp_bm8_int8; -pub mod mpp_bm64; -pub mod mpp_bm64_int8; -pub mod mpp_block_scaled; -pub mod mpp_bm8_block_scaled; -pub mod mpp_bm64_block_scaled; +// test/bench helpers in `moe_mpp_shared`). +pub mod moe_mpp; +pub mod moe_mpp_shared; +pub mod moe_mpp_int8; +pub mod moe_mpp_bm8; +pub mod moe_mpp_bm8_int8; +pub mod moe_mpp_bm64; +pub mod moe_mpp_bm64_int8; +pub mod moe_mpp_block_scaled; +pub mod moe_mpp_bm8_block_scaled; +pub mod moe_mpp_bm64_block_scaled; // GGUF-format per-expert matmul / matvec (q2k, iq2xxs, q4). -pub mod bgemm_q2k_bm64; -pub mod bgemm_q2k_mpp; -pub mod bgemm_q2k_view; -pub mod bgemm_q2k_view_u16_bm64; -pub mod bgemm_q4_bm64; -pub mod bgemm_iq2xxs_bm64; -pub mod bgemm_iq2xxs_mpp; -pub mod bgemm_iq2xxs_view; -pub mod bgemm_iq2xxs_view_u16_bm64; -pub mod gather_down_q2k; -pub mod gather_gemv_iq2xxs; -pub mod gemv_rows_q2k; -pub mod gemv_rows_iq2xxs; -pub mod gemv_rows_view_iq2xxs; -pub mod gemv_ws_q2k; -pub mod gemv_ws_iq2xxs; +pub mod moe_bgemm_q2k_bm64; +pub mod moe_bgemm_q2k_mpp; +pub mod moe_bgemm_q2k_view; +pub mod moe_bgemm_q2k_view_u16_bm64; +pub mod moe_bgemm_q4_bm64; +pub mod moe_bgemm_iq2xxs_bm64; +pub mod moe_bgemm_iq2xxs_mpp; +pub mod moe_bgemm_iq2xxs_view; +pub mod moe_bgemm_iq2xxs_view_u16_bm64; +pub mod moe_gather_down_q2k; +pub mod moe_gather_gemv_iq2xxs; +pub mod moe_gemv_rows_q2k; +pub mod moe_gemv_rows_iq2xxs; +pub mod moe_gemv_rows_view_iq2xxs; +pub mod moe_gemv_ws_q2k; +pub mod moe_gemv_ws_iq2xxs; // Batched Q4 expert gather (up / down / weighted-sum), seeded from gemv_q8. -pub mod gather_q4; +pub mod moe_gather_q4; // Down-projection combine (swiglu-fused accumulate, weighted sum). -pub mod down_swiglu_accum; -pub mod down_weighted_sum_f16; +pub mod moe_down_swiglu_accum; +pub mod moe_down_weighted_sum_f16; // Expert-indexed dequant GEMV + block-scaled MoE matmul. pub mod dequant_gemv_expert_indexed; pub mod dequant_gemv_expert_indexed_block_scaled; -pub mod block_scaled; +pub mod block_scaled_moe; diff --git a/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_bm64.rs b/crates/metaltile-std/src/kernels/moe/moe_bgemm_iq2xxs_bm64.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_bm64.rs rename to crates/metaltile-std/src/kernels/moe/moe_bgemm_iq2xxs_bm64.rs diff --git a/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_mpp.rs b/crates/metaltile-std/src/kernels/moe/moe_bgemm_iq2xxs_mpp.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_mpp.rs rename to crates/metaltile-std/src/kernels/moe/moe_bgemm_iq2xxs_mpp.rs diff --git a/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view.rs b/crates/metaltile-std/src/kernels/moe/moe_bgemm_iq2xxs_view.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view.rs rename to crates/metaltile-std/src/kernels/moe/moe_bgemm_iq2xxs_view.rs diff --git a/crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view_u16_bm64.rs b/crates/metaltile-std/src/kernels/moe/moe_bgemm_iq2xxs_view_u16_bm64.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/bgemm_iq2xxs_view_u16_bm64.rs rename to crates/metaltile-std/src/kernels/moe/moe_bgemm_iq2xxs_view_u16_bm64.rs diff --git a/crates/metaltile-std/src/kernels/moe/bgemm_q2k_bm64.rs b/crates/metaltile-std/src/kernels/moe/moe_bgemm_q2k_bm64.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/bgemm_q2k_bm64.rs rename to crates/metaltile-std/src/kernels/moe/moe_bgemm_q2k_bm64.rs diff --git a/crates/metaltile-std/src/kernels/moe/bgemm_q2k_mpp.rs b/crates/metaltile-std/src/kernels/moe/moe_bgemm_q2k_mpp.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/bgemm_q2k_mpp.rs rename to crates/metaltile-std/src/kernels/moe/moe_bgemm_q2k_mpp.rs diff --git a/crates/metaltile-std/src/kernels/moe/bgemm_q2k_view.rs b/crates/metaltile-std/src/kernels/moe/moe_bgemm_q2k_view.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/bgemm_q2k_view.rs rename to crates/metaltile-std/src/kernels/moe/moe_bgemm_q2k_view.rs diff --git a/crates/metaltile-std/src/kernels/moe/bgemm_q2k_view_u16_bm64.rs b/crates/metaltile-std/src/kernels/moe/moe_bgemm_q2k_view_u16_bm64.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/bgemm_q2k_view_u16_bm64.rs rename to crates/metaltile-std/src/kernels/moe/moe_bgemm_q2k_view_u16_bm64.rs diff --git a/crates/metaltile-std/src/kernels/moe/bgemm_q4_bm64.rs b/crates/metaltile-std/src/kernels/moe/moe_bgemm_q4_bm64.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/bgemm_q4_bm64.rs rename to crates/metaltile-std/src/kernels/moe/moe_bgemm_q4_bm64.rs diff --git a/crates/metaltile-std/src/kernels/moe/down_swiglu_accum.rs b/crates/metaltile-std/src/kernels/moe/moe_down_swiglu_accum.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/down_swiglu_accum.rs rename to crates/metaltile-std/src/kernels/moe/moe_down_swiglu_accum.rs diff --git a/crates/metaltile-std/src/kernels/moe/down_weighted_sum_f16.rs b/crates/metaltile-std/src/kernels/moe/moe_down_weighted_sum_f16.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/down_weighted_sum_f16.rs rename to crates/metaltile-std/src/kernels/moe/moe_down_weighted_sum_f16.rs diff --git a/crates/metaltile-std/src/kernels/moe/gather_down_q2k.rs b/crates/metaltile-std/src/kernels/moe/moe_gather_down_q2k.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/gather_down_q2k.rs rename to crates/metaltile-std/src/kernels/moe/moe_gather_down_q2k.rs diff --git a/crates/metaltile-std/src/kernels/moe/gather_gemv_iq2xxs.rs b/crates/metaltile-std/src/kernels/moe/moe_gather_gemv_iq2xxs.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/gather_gemv_iq2xxs.rs rename to crates/metaltile-std/src/kernels/moe/moe_gather_gemv_iq2xxs.rs diff --git a/crates/metaltile-std/src/kernels/moe/gather_q4.rs b/crates/metaltile-std/src/kernels/moe/moe_gather_q4.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/gather_q4.rs rename to crates/metaltile-std/src/kernels/moe/moe_gather_q4.rs diff --git a/crates/metaltile-std/src/kernels/moe/gather_qmm.rs b/crates/metaltile-std/src/kernels/moe/moe_gather_qmm.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/gather_qmm.rs rename to crates/metaltile-std/src/kernels/moe/moe_gather_qmm.rs diff --git a/crates/metaltile-std/src/kernels/moe/gemv_rows_iq2xxs.rs b/crates/metaltile-std/src/kernels/moe/moe_gemv_rows_iq2xxs.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/gemv_rows_iq2xxs.rs rename to crates/metaltile-std/src/kernels/moe/moe_gemv_rows_iq2xxs.rs diff --git a/crates/metaltile-std/src/kernels/moe/gemv_rows_q2k.rs b/crates/metaltile-std/src/kernels/moe/moe_gemv_rows_q2k.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/gemv_rows_q2k.rs rename to crates/metaltile-std/src/kernels/moe/moe_gemv_rows_q2k.rs diff --git a/crates/metaltile-std/src/kernels/moe/gemv_rows_view_iq2xxs.rs b/crates/metaltile-std/src/kernels/moe/moe_gemv_rows_view_iq2xxs.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/gemv_rows_view_iq2xxs.rs rename to crates/metaltile-std/src/kernels/moe/moe_gemv_rows_view_iq2xxs.rs diff --git a/crates/metaltile-std/src/kernels/moe/gemv_ws_iq2xxs.rs b/crates/metaltile-std/src/kernels/moe/moe_gemv_ws_iq2xxs.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/gemv_ws_iq2xxs.rs rename to crates/metaltile-std/src/kernels/moe/moe_gemv_ws_iq2xxs.rs diff --git a/crates/metaltile-std/src/kernels/moe/gemv_ws_q2k.rs b/crates/metaltile-std/src/kernels/moe/moe_gemv_ws_q2k.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/gemv_ws_q2k.rs rename to crates/metaltile-std/src/kernels/moe/moe_gemv_ws_q2k.rs diff --git a/crates/metaltile-std/src/kernels/moe/mpp.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp.rs similarity index 98% rename from crates/metaltile-std/src/kernels/moe/mpp.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp.rs index 43f0741d..e172cca4 100644 --- a/crates/metaltile-std/src/kernels/moe/mpp.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp.rs @@ -250,7 +250,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int4_bm16_mpp; - use crate::kernels::moe::mpp_shared::{MmaTestShape, int4_indexed_setup}; + use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int4_bm16_mpp(dt: DType) -> TestSetup { @@ -271,7 +271,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int4_bm16_mpp; - use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int4_bm16_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/kernels/moe/mpp_block_scaled.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_block_scaled.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/mpp_block_scaled.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp_block_scaled.rs diff --git a/crates/metaltile-std/src/kernels/moe/mpp_bm64.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs similarity index 98% rename from crates/metaltile-std/src/kernels/moe/mpp_bm64.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs index f542d03a..17d095f4 100644 --- a/crates/metaltile-std/src/kernels/moe/mpp_bm64.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs @@ -220,7 +220,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int4_bm64_mpp; - use crate::kernels::moe::mpp_shared::{MmaTestShape, int4_indexed_setup}; + use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int4_bm64_mpp(dt: DType) -> TestSetup { @@ -241,7 +241,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int4_bm64_mpp; - use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int4_bm64_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/kernels/moe/mpp_bm64_block_scaled.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64_block_scaled.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/mpp_bm64_block_scaled.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp_bm64_block_scaled.rs diff --git a/crates/metaltile-std/src/kernels/moe/mpp_bm64_int8.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64_int8.rs similarity index 98% rename from crates/metaltile-std/src/kernels/moe/mpp_bm64_int8.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp_bm64_int8.rs index b7605362..7691f903 100644 --- a/crates/metaltile-std/src/kernels/moe/mpp_bm64_int8.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64_int8.rs @@ -245,7 +245,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int8_bm64_mpp; - use crate::kernels::moe::mpp_shared::{MmaTestShape, int8_indexed_setup}; + use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int8_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int8_bm64_mpp(dt: DType) -> TestSetup { @@ -270,7 +270,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int8_bm64_mpp; - use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int8_bm64_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/kernels/moe/mpp_bm8.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs similarity index 98% rename from crates/metaltile-std/src/kernels/moe/mpp_bm8.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs index 80fc9e6e..ff6f8af4 100644 --- a/crates/metaltile-std/src/kernels/moe/mpp_bm8.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs @@ -214,7 +214,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int4_bm8_mpp; - use crate::kernels::moe::mpp_shared::{MmaTestShape, int4_indexed_setup}; + use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int4_bm8_mpp(dt: DType) -> TestSetup { @@ -235,7 +235,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int4_bm8_mpp; - use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int4_bm8_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/kernels/moe/mpp_bm8_block_scaled.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8_block_scaled.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/mpp_bm8_block_scaled.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp_bm8_block_scaled.rs diff --git a/crates/metaltile-std/src/kernels/moe/mpp_bm8_int8.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8_int8.rs similarity index 98% rename from crates/metaltile-std/src/kernels/moe/mpp_bm8_int8.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp_bm8_int8.rs index 05ea0351..f5a10ad7 100644 --- a/crates/metaltile-std/src/kernels/moe/mpp_bm8_int8.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8_int8.rs @@ -241,7 +241,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int8_bm8_mpp; - use crate::kernels::moe::mpp_shared::{MmaTestShape, int8_indexed_setup}; + use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int8_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int8_bm8_mpp(dt: DType) -> TestSetup { @@ -265,7 +265,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int8_bm8_mpp; - use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int8_bm8_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/kernels/moe/mpp_int8.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_int8.rs similarity index 98% rename from crates/metaltile-std/src/kernels/moe/mpp_int8.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp_int8.rs index 658afc8d..d4d2ba5d 100644 --- a/crates/metaltile-std/src/kernels/moe/mpp_int8.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_int8.rs @@ -257,7 +257,7 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::mt_moe_gather_qmm_mma_int8_bm16_mpp; - use crate::kernels::moe::mpp_shared::{MmaTestShape, int8_indexed_setup}; + use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int8_indexed_setup}; #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int8_bm16_mpp(dt: DType) -> TestSetup { @@ -281,7 +281,7 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::mt_moe_gather_qmm_mma_int8_bm16_mpp; - use crate::kernels::moe::mpp_shared::{MmaBenchShape, int4_mma_bench}; + use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] fn bench_moe_gather_qmm_mma_int8_bm16_mpp(dt: DType) -> BenchSetup { diff --git a/crates/metaltile-std/src/kernels/moe/mpp_shared.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_shared.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/mpp_shared.rs rename to crates/metaltile-std/src/kernels/moe/moe_mpp_shared.rs diff --git a/crates/metaltile-std/src/kernels/moe/permute.rs b/crates/metaltile-std/src/kernels/moe/moe_permute.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/permute.rs rename to crates/metaltile-std/src/kernels/moe/moe_permute.rs diff --git a/crates/metaltile-std/src/kernels/moe/router_sigmoid_bias.rs b/crates/metaltile-std/src/kernels/moe/moe_router_sigmoid_bias.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/router_sigmoid_bias.rs rename to crates/metaltile-std/src/kernels/moe/moe_router_sigmoid_bias.rs diff --git a/crates/metaltile-std/src/kernels/moe/router_sqrtsoftplus.rs b/crates/metaltile-std/src/kernels/moe/moe_router_sqrtsoftplus.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/router_sqrtsoftplus.rs rename to crates/metaltile-std/src/kernels/moe/moe_router_sqrtsoftplus.rs diff --git a/crates/metaltile-std/src/kernels/moe/router_topk.rs b/crates/metaltile-std/src/kernels/moe/moe_router_topk.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/router_topk.rs rename to crates/metaltile-std/src/kernels/moe/moe_router_topk.rs diff --git a/crates/metaltile-std/src/kernels/moe/router_topk_biased.rs b/crates/metaltile-std/src/kernels/moe/moe_router_topk_biased.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/router_topk_biased.rs rename to crates/metaltile-std/src/kernels/moe/moe_router_topk_biased.rs diff --git a/crates/metaltile-std/src/kernels/moe/sigmoid_bias.rs b/crates/metaltile-std/src/kernels/moe/moe_sigmoid_bias.rs similarity index 100% rename from crates/metaltile-std/src/kernels/moe/sigmoid_bias.rs rename to crates/metaltile-std/src/kernels/moe/moe_sigmoid_bias.rs diff --git a/crates/metaltile-std/tests/bm64_vs_gemvrows_iq2xxs.rs b/crates/metaltile-std/tests/bm64_vs_gemvrows_iq2xxs.rs index 6e9faac5..fbb43f25 100644 --- a/crates/metaltile-std/tests/bm64_vs_gemvrows_iq2xxs.rs +++ b/crates/metaltile-std/tests/bm64_vs_gemvrows_iq2xxs.rs @@ -13,8 +13,8 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; use metaltile_std::kernels::moe::{ - bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, - gemv_rows_iq2xxs::mt_moe_gemv_rows_iq2xxs, + moe_bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, + moe_gemv_rows_iq2xxs::mt_moe_gemv_rows_iq2xxs, }; fn xorshift(s: &mut u32) -> u32 { diff --git a/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs b/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs index c7b80f0a..a133ea47 100644 --- a/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs +++ b/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs @@ -10,7 +10,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes, unpack_u32_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::router_topk_biased::{mt_moe_router_topk_biased, mt_remap_u32}; +use metaltile_std::kernels::moe::moe_router_topk_biased::{mt_moe_router_topk_biased, mt_remap_u32}; #[test] fn dsv4_router_topk_f32() { diff --git a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_bm64_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_bm64_correctness.rs index 8eb6bb66..bf371d6a 100644 --- a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_bm64_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_bm64_correctness.rs @@ -13,8 +13,8 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; use metaltile_std::kernels::moe::{ - bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, - bgemm_iq2xxs_mpp::mt_moe_gather_bgemm_iq2xxs_mpp, + moe_bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, + moe_bgemm_iq2xxs_mpp::mt_moe_gather_bgemm_iq2xxs_mpp, }; fn xorshift(s: &mut u32) -> u32 { diff --git a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_mpp_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_mpp_correctness.rs index 775339e1..ed09e370 100644 --- a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_mpp_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_mpp_correctness.rs @@ -12,7 +12,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::bgemm_iq2xxs_mpp::mt_moe_gather_bgemm_iq2xxs_mpp; +use metaltile_std::kernels::moe::moe_bgemm_iq2xxs_mpp::mt_moe_gather_bgemm_iq2xxs_mpp; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; diff --git a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_view_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_view_correctness.rs index a1cd0a8e..85952696 100644 --- a/crates/metaltile-std/tests/moe_bgemm_iq2xxs_view_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_iq2xxs_view_correctness.rs @@ -19,7 +19,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use half::f16; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::bgemm_iq2xxs_view::mt_moe_bgemm_iq2xxs_view; +use metaltile_std::kernels::moe::moe_bgemm_iq2xxs_view::mt_moe_bgemm_iq2xxs_view; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; diff --git a/crates/metaltile-std/tests/moe_bgemm_q2k_bm64_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_q2k_bm64_correctness.rs index 7a7eef9d..224809fb 100644 --- a/crates/metaltile-std/tests/moe_bgemm_q2k_bm64_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_q2k_bm64_correctness.rs @@ -11,8 +11,8 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; use metaltile_std::kernels::moe::{ - bgemm_q2k_bm64::mt_moe_bgemm_q2k_bm64, - bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp, + moe_bgemm_q2k_bm64::mt_moe_bgemm_q2k_bm64, + moe_bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp, }; fn xorshift(s: &mut u32) -> u32 { diff --git a/crates/metaltile-std/tests/moe_bgemm_q2k_mpp_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_q2k_mpp_correctness.rs index 5d7119b6..abf0a822 100644 --- a/crates/metaltile-std/tests/moe_bgemm_q2k_mpp_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_q2k_mpp_correctness.rs @@ -10,7 +10,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp; +use metaltile_std::kernels::moe::moe_bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp; // Shared Q2_K output-index → (qs byte, 2-bit shift) map (see PR #264/#265): the // kernel, quantizer, and this oracle all read the one definition in quant::gguf. use metaltile_std::quant::gguf::q2_k_qpos; diff --git a/crates/metaltile-std/tests/moe_bgemm_q2k_view_correctness.rs b/crates/metaltile-std/tests/moe_bgemm_q2k_view_correctness.rs index 46b0979d..df5789a8 100644 --- a/crates/metaltile-std/tests/moe_bgemm_q2k_view_correctness.rs +++ b/crates/metaltile-std/tests/moe_bgemm_q2k_view_correctness.rs @@ -17,8 +17,8 @@ use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use half::f16; use metaltile::{Context, core::ir::KernelMode}; use metaltile_std::kernels::moe::{ - bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp, - bgemm_q2k_view::mt_moe_bgemm_q2k_view, + moe_bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp, + moe_bgemm_q2k_view::mt_moe_bgemm_q2k_view, }; fn xorshift(s: &mut u32) -> u32 { diff --git a/crates/metaltile-std/tests/moe_bm64_ragged_correctness.rs b/crates/metaltile-std/tests/moe_bm64_ragged_correctness.rs index b3c42ca3..c1e235c4 100644 --- a/crates/metaltile-std/tests/moe_bm64_ragged_correctness.rs +++ b/crates/metaltile-std/tests/moe_bm64_ragged_correctness.rs @@ -22,8 +22,8 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; use metaltile_std::kernels::moe::{ - bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, - gemv_rows_iq2xxs::mt_moe_gemv_rows_iq2xxs, + moe_bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, + moe_gemv_rows_iq2xxs::mt_moe_gemv_rows_iq2xxs, }; fn read_u32(p: &str) -> Vec { diff --git a/crates/metaltile-std/tests/moe_gather_down_q2k_correctness.rs b/crates/metaltile-std/tests/moe_gather_down_q2k_correctness.rs index 1119f634..607a85a9 100644 --- a/crates/metaltile-std/tests/moe_gather_down_q2k_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_down_q2k_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `kernels::moe::gather_down_q2k` — fused 6-expert +//! GPU correctness for `kernels::moe::moe_gather_down_q2k` — fused 6-expert //! Q2_K inline-dequant down-projection + router-weighted sum. Validates //! against a CPU reference running the identical (production-proven) //! Q2_K dequant formula. @@ -15,7 +15,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::gather_down_q2k::mt_moe_gather_down_q2k; +use metaltile_std::kernels::moe::moe_gather_down_q2k::mt_moe_gather_down_q2k; // The Q2_K output-index → (qs byte, 2-bit shift) map is the single shared // definition in `quant::gguf`: the kernel, the quantizer, and this oracle all // read it, so the layout can't drift apart (getting it wrong was PR #264). diff --git a/crates/metaltile-std/tests/moe_gather_gemv_iq2xxs_correctness.rs b/crates/metaltile-std/tests/moe_gather_gemv_iq2xxs_correctness.rs index 972ef2ea..5f0441c1 100644 --- a/crates/metaltile-std/tests/moe_gather_gemv_iq2xxs_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_gemv_iq2xxs_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `kernels::moe::gather_gemv_iq2xxs` — the fused +//! GPU correctness for `kernels::moe::moe_gather_gemv_iq2xxs` — the fused //! 6-expert IQ2_XXS inline-dequant gather GEMV used by the DSv4 decode //! FFN. Validates GPU output against a CPU reference that runs the //! identical (production-proven) IQ2_XXS dequant formula, so a wrong @@ -17,7 +17,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::gather_gemv_iq2xxs::mt_moe_gather_gemv_iq2xxs; +use metaltile_std::kernels::moe::moe_gather_gemv_iq2xxs::mt_moe_gather_gemv_iq2xxs; const N_SLOTS: usize = 6; diff --git a/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs index 2b82d63b..6b33fd46 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_int4_m16_m32_correctness.rs @@ -24,7 +24,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::gather_qmm::{ +use metaltile_std::kernels::moe::moe_gather_qmm::{ mt_moe_gather_qmm_int4_m8, mt_moe_gather_qmm_int4_m16, mt_moe_gather_qmm_int4_m32, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs b/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs index 4625ff0d..b3c64d67 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_microbench.rs @@ -13,7 +13,7 @@ use std::{collections::BTreeMap, time::Instant}; use common::{Dt, gpu_lock, pack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::gather_qmm::{ +use metaltile_std::kernels::moe::moe_gather_qmm::{ mt_moe_gather_qmm_int4, mt_moe_gather_qmm_int4_m8, mt_moe_gather_qmm_mma_int4, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs index e4056cc4..a70d81f0 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mma_bitwidth_correctness.rs @@ -3,7 +3,7 @@ #![allow(clippy::manual_is_multiple_of)] //! GPU correctness for the bit-width-generalized MMA MoE BGEMMs -//! `kernels::moe::gather_qmm::mt_moe_gather_qmm_mma_b{3,5,6,8}`. +//! `kernels::moe::moe_gather_qmm::mt_moe_gather_qmm_mma_b{3,5,6,8}`. //! //! Same tiled-MMA algorithm as `mt_moe_gather_qmm_mma_int4`, but the //! weight coop-dequant pulls codes from a contiguous LSB-first bit-stream @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::gather_qmm::{ +use metaltile_std::kernels::moe::moe_gather_qmm::{ mt_moe_gather_qmm_mma_b3, mt_moe_gather_qmm_mma_b5, mt_moe_gather_qmm_mma_b6, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs index e1a82ad3..b390c78f 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `kernels::moe::mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp`. +//! GPU correctness for `kernels::moe::moe_mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp`. //! //! BM=BN=64 MPP MoE kernel — same output semantics as the BM=16 sibling but //! scaled up to a 64×64 output tile with 4 SGs (WM=WN=2) per TG. Validated @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_int4, mpp_bm64}; +use metaltile_std::kernels::moe::{moe_gather_qmm::mt_moe_gather_qmm_int4, moe_mpp_bm64}; /// Pack a row of int4 weights into uint32s (8 per uint, LSB-first per /// nibble). Identical to the helper used by the bm16_mpp test — @@ -148,7 +148,7 @@ fn moe_gather_qmm_mma_int4_bm64_mpp_matches_m1_clean_tile() { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); + moe_mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [ceil(N/64), ceil(T/64), 1]. TG: 128 lanes = 4 SGs (WM=WN=2). let r = ctx @@ -267,7 +267,7 @@ fn moe_gather_qmm_mma_int4_bm64_mpp_matches_m1_multi_tile() { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); + moe_mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid( @@ -391,7 +391,7 @@ fn moe_gather_qmm_mma_int4_bm64_mpp_bf16_matches_m1_clean_tile() { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::Bf16.to_dtype()); + moe_mpp_bm64::mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(Dt::Bf16.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid( diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs index 4efc3bad..2a54370b 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `kernels::moe::mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp`. +//! GPU correctness for `kernels::moe::moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp`. //! //! BM=BN=64 MPP MoE int8 kernel — same output semantics as the int4 BM=64 //! sibling but the weight layout changes from 8 nibbles/u32 to 4 bytes/u32. @@ -21,7 +21,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_b8, mpp_bm64_int8}; +use metaltile_std::kernels::moe::{moe_gather_qmm::mt_moe_gather_qmm_b8, moe_mpp_bm64_int8}; /// Pack a row of int8 weight codes into uint32s (4 codes per uint, LE byte /// order). Code values must be in 0..=255. @@ -134,7 +134,7 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_matches_b8_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( + let mut k = moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( Dt::F32.to_dtype(), ); k.mode = KernelMode::Reduction; @@ -250,7 +250,7 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_matches_b8_multi_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( + let mut k = moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( Dt::F32.to_dtype(), ); k.mode = KernelMode::Reduction; @@ -370,7 +370,7 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_bf16_matches_b8_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( + let mut k = moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( Dt::Bf16.to_dtype(), ); k.mode = KernelMode::Reduction; diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs index bbff0b09..da7b7248 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `kernels::moe::mpp_bm8::mt_moe_gather_qmm_mma_int4_bm8_mpp`. +//! GPU correctness for `kernels::moe::moe_mpp_bm8::mt_moe_gather_qmm_mma_int4_bm8_mpp`. //! //! BM=8 MPP MoE kernel — same output semantics as the BM=16 / BM=64 siblings //! but the per-TG row tile shrinks to 8 to match decode-time MoE shapes @@ -23,7 +23,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_int4, mpp_bm8}; +use metaltile_std::kernels::moe::{moe_gather_qmm::mt_moe_gather_qmm_int4, moe_mpp_bm8}; /// Pack a row of int4 weights into uint32s (8 per uint, LSB-first per /// nibble). Same helper used by the bm16_mpp / bm64_mpp test files — @@ -187,7 +187,7 @@ fn run_case(case: &Case) { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = mpp_bm8::mt_moe_gather_qmm_mma_int4_bm8_mpp::kernel_ir_for(dt.to_dtype()); + let mut k = moe_mpp_bm8::mt_moe_gather_qmm_mma_int4_bm8_mpp::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [ceil(N/32), ceil(T/8), 1]. TG: 32 lanes = 1 SG. let r = ctx diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs index a6be61be..2b1ff04c 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `kernels::moe::mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp`. +//! GPU correctness for `kernels::moe::moe_mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp`. //! //! BM=8 MPP MoE int8 kernel — same output semantics as the int4 BM=8 sibling //! but the weight layout changes from 8 nibbles/u32 to 4 bytes/u32. @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_b8, mpp_bm8_int8}; +use metaltile_std::kernels::moe::{moe_gather_qmm::mt_moe_gather_qmm_b8, moe_mpp_bm8_int8}; /// Pack a row of int8 weight codes into uint32s (4 codes per uint, LE byte /// order). Code values must be in 0..=255. @@ -174,7 +174,7 @@ fn run_case(case: &Case) { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt.to_dtype()); + moe_mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [ceil(N/32), ceil(T/8), 1]. TG: 32 lanes = 1 SG. let r = ctx diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs index e4e8f18f..1fe7893c 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `kernels::moe::mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp`. +//! GPU correctness for `kernels::moe::moe_mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp`. //! //! This is the MPP (MetalPerformancePrimitives) MoE BGEMM — same algorithm //! and output as `mt_moe_gather_qmm_mma_int4_bm16` but routes the inner @@ -23,7 +23,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{gather_qmm::mt_moe_gather_qmm_int4, mpp}; +use metaltile_std::kernels::moe::{moe_gather_qmm::mt_moe_gather_qmm_int4, moe_mpp}; /// Pack a row of int4 weights into uint32s (8 per uint, LSB-first per nibble). /// Identical to the helper used by the legacy @@ -141,7 +141,7 @@ fn moe_gather_qmm_mma_int4_bm16_mpp_matches_m1_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(Dt::F32.to_dtype()); + let mut k = moe_mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [N/BN=32, ceil(T/BM=16), 1]. TG: 32 lanes = 1 SG (MPP's // matmul2d uses `execution_simdgroup`). @@ -274,7 +274,7 @@ fn moe_gather_qmm_mma_int4_bm16_mpp_bf16_matches_m1_clean_tile() { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(Dt::Bf16.to_dtype()); + moe_mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(Dt::Bf16.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid( diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs index 7341c920..45ff3364 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `kernels::moe::mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp`. +//! GPU correctness for `kernels::moe::moe_mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp`. //! //! MPP (MetalPerformancePrimitives) int8 MoE BGEMM — same algorithm as //! `mt_moe_gather_qmm_mma_int4_bm16_mpp` but with pack-aligned 8-bit @@ -21,7 +21,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::mpp_int8; +use metaltile_std::kernels::moe::moe_mpp_int8; // ── helpers ──────────────────────────────────────────────────────────────── @@ -132,7 +132,7 @@ fn run_mpp_int8( buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("Context::new"); - let mut k = mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt.to_dtype()); + let mut k = moe_mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [N/BN=32, ceil(T/BM=16), 1], TG: [32, 1, 1] (1 SG — MPP matmul2d). let r = ctx diff --git a/crates/metaltile-std/tests/moe_gemv_rows_iq2xxs_correctness.rs b/crates/metaltile-std/tests/moe_gemv_rows_iq2xxs_correctness.rs index 4d55df72..b2a2f3f1 100644 --- a/crates/metaltile-std/tests/moe_gemv_rows_iq2xxs_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_rows_iq2xxs_correctness.rs @@ -12,7 +12,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::gemv_rows_iq2xxs::mt_moe_gemv_rows_iq2xxs; +use metaltile_std::kernels::moe::moe_gemv_rows_iq2xxs::mt_moe_gemv_rows_iq2xxs; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; diff --git a/crates/metaltile-std/tests/moe_gemv_rows_q2k_correctness.rs b/crates/metaltile-std/tests/moe_gemv_rows_q2k_correctness.rs index 3068fe8b..af73b427 100644 --- a/crates/metaltile-std/tests/moe_gemv_rows_q2k_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_rows_q2k_correctness.rs @@ -13,8 +13,8 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; use metaltile_std::kernels::moe::{ - bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp, - gemv_rows_q2k::mt_moe_gemv_rows_q2k, + moe_bgemm_q2k_mpp::mt_moe_gather_bgemm_q2k_mpp, + moe_gemv_rows_q2k::mt_moe_gemv_rows_q2k, }; fn xorshift(s: &mut u32) -> u32 { diff --git a/crates/metaltile-std/tests/moe_gemv_rows_view_iq2xxs_correctness.rs b/crates/metaltile-std/tests/moe_gemv_rows_view_iq2xxs_correctness.rs index 1b4b227d..ce233ffc 100644 --- a/crates/metaltile-std/tests/moe_gemv_rows_view_iq2xxs_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_rows_view_iq2xxs_correctness.rs @@ -16,7 +16,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use half::f16; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::gemv_rows_view_iq2xxs::{ +use metaltile_std::kernels::moe::moe_gemv_rows_view_iq2xxs::{ mt_moe_gemv_rows_view_iq2xxs, mt_moe_gemv_rows_view_u16_iq2xxs, }; diff --git a/crates/metaltile-std/tests/moe_gemv_ws_iq2xxs_correctness.rs b/crates/metaltile-std/tests/moe_gemv_ws_iq2xxs_correctness.rs index 482283a5..fb73021b 100644 --- a/crates/metaltile-std/tests/moe_gemv_ws_iq2xxs_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_ws_iq2xxs_correctness.rs @@ -14,7 +14,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::gemv_ws_iq2xxs::mt_moe_gemv_ws_iq2xxs; +use metaltile_std::kernels::moe::moe_gemv_ws_iq2xxs::mt_moe_gemv_ws_iq2xxs; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; diff --git a/crates/metaltile-std/tests/moe_gemv_ws_q2k_correctness.rs b/crates/metaltile-std/tests/moe_gemv_ws_q2k_correctness.rs index e99f7057..972bbda1 100644 --- a/crates/metaltile-std/tests/moe_gemv_ws_q2k_correctness.rs +++ b/crates/metaltile-std/tests/moe_gemv_ws_q2k_correctness.rs @@ -15,8 +15,8 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, pack_u32_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; use metaltile_std::kernels::moe::{ - gemv_rows_q2k::mt_moe_gemv_rows_q2k, - gemv_ws_q2k::mt_moe_gemv_ws_q2k, + moe_gemv_rows_q2k::mt_moe_gemv_rows_q2k, + moe_gemv_ws_q2k::mt_moe_gemv_ws_q2k, }; fn xorshift(s: &mut u32) -> u32 { diff --git a/crates/metaltile-std/tests/moe_q2k_view_u16_correctness.rs b/crates/metaltile-std/tests/moe_q2k_view_u16_correctness.rs index 085ec2bf..7f25b4af 100644 --- a/crates/metaltile-std/tests/moe_q2k_view_u16_correctness.rs +++ b/crates/metaltile-std/tests/moe_q2k_view_u16_correctness.rs @@ -13,8 +13,8 @@ use metaltile::{ core::{dtype::DType, ir::KernelMode}, }; use metaltile_std::kernels::moe::{ - bgemm_q2k_bm64::mt_moe_bgemm_q2k_bm64, - bgemm_q2k_view_u16_bm64::mt_moe_bgemm_q2k_view_u16_bm64, + moe_bgemm_q2k_bm64::mt_moe_bgemm_q2k_bm64, + moe_bgemm_q2k_view_u16_bm64::mt_moe_bgemm_q2k_view_u16_bm64, }; fn xs(s: &mut u32) -> u32 { diff --git a/crates/metaltile-std/tests/moe_view_u16_correctness.rs b/crates/metaltile-std/tests/moe_view_u16_correctness.rs index 3fc7b4a7..4387d092 100644 --- a/crates/metaltile-std/tests/moe_view_u16_correctness.rs +++ b/crates/metaltile-std/tests/moe_view_u16_correctness.rs @@ -17,8 +17,8 @@ use metaltile::{ core::{dtype::DType, ir::KernelMode}, }; use metaltile_std::kernels::moe::{ - bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, - bgemm_iq2xxs_view_u16_bm64::mt_moe_bgemm_iq2xxs_view_u16_bm64, + moe_bgemm_iq2xxs_bm64::mt_moe_bgemm_iq2xxs_bm64, + moe_bgemm_iq2xxs_view_u16_bm64::mt_moe_bgemm_iq2xxs_view_u16_bm64, }; struct Lcg(u64); diff --git a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md index 57928528..4e210cbd 100644 --- a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md +++ b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md @@ -59,7 +59,7 @@ crates/metaltile-std/src/kernels/ router_topk_biased / sigmoid_bias / sqrtsoftplus · mpp(bm8/bm64 × int8 × block_scaled) + mpp_shared · bgemm/gemv (q2k/iq2xxs/q4, view/ws/rows) · gather_q4 · down_swiglu_accum / down_weighted_sum · dequant_gemv_expert_indexed(_block_scaled) · - block_scaled. Filenames drop the moe_ prefix; format-axis fold deferred (§7). + block_scaled. Filenames keep the moe_ prefix (match kernel names); format-axis fold deferred (§7). (orchestration split into router_topk / permute / gather_qmm.) norm/ ✅ DONE — rms_norm(+residual/rope/qgemv/gated) · layer_norm · adain1d rope/ ✅ DONE — rope · rope_2d · rope_banded · rope_yarn · partial_rope From 1f440450e4e62a4388b9e4446b81ee2f55cf4713 Mon Sep 17 00:00:00 2001 From: Eric Kryski <599019+ekryski@users.noreply.github.com> Date: Sat, 13 Jun 2026 14:18:58 -0600 Subject: [PATCH 4/7] refactor(moe): fold int4/int8 mpp BGEMM onto a BITS variant axis (6->3 files) The six integer MPP grouped-BGEMM files were three {int4,int8} pairs that share an identical matmul path (same coop_tile descriptor / staging / write-back) and differ only in the weight unpack. Fold each pair onto a compile-time BITS axis: moe_mpp.rs = variants(BITS=[4,8]) bm16 (cooperative-tensor path) moe_mpp_bm8.rs = variants(BITS=[4,8]) bm8 (direct-input path) moe_mpp_bm64.rs = variants(BITS=[4,8]) bm64 (4-simdgroup 2x2 path) The unpack parameterizes cleanly with no branch: vals_per_pack = 32/BITS codes per u32, decoded by (packed >> j*BITS) & ((1<> (j*BITS)) & ((1<( +pub fn mt_moe_gather_qmm_mma( x: Tensor, w: Tensor, scales: Tensor, @@ -65,7 +60,9 @@ pub fn mt_moe_gather_qmm_mma_int4_bm16_mpp( let n_tile_base = tgid_x * 32u32; let m_tile_base = tgid_y * 16u32; let lane = simd_lane; - let packs_per_row = k_in / 8u32; + // Weight packing: `32/BITS` codes per u32 (int4 → 8 nibbles, int8 → 4 bytes). + let vals_per_pack = 32u32 / BITS; + let packs_per_row = k_in / vals_per_pack; let groups_per_row = k_in / group_size; // Threadgroup staging tiles. `coop_stage(T)` = half for bf16, else T — // the matmul reads these as cooperative tensors. `out_scratch` is @@ -131,24 +128,26 @@ pub fn mt_moe_gather_qmm_mma_int4_bm16_mpp( threadgroup_store("xs", mr * 16u32 + kc, select(in_run, xv, 0.0f32)); } // Dequant W[expert, n_tile_base..+32, kb..kb+16] → ws. - // 32 lanes × 2 packs/lane; 8 nibbles/pack. - for _pi in range(0u32, 2u32, 1u32) { - let pack_id = lane * 2u32 + _pi; - let w_row = pack_id / 2u32; // 0..31 (BN rows) - let pack_col = pack_id % 2u32; // 0..1 (BK=16 → 2 packs) + // 32 lanes × `packs_per_lane` packs/lane; `vals_per_pack` codes/pack. + let packs_per_lane = 16u32 / vals_per_pack; + let mask = (1u32 << BITS) - 1u32; + for _pi in range(0u32, packs_per_lane, 1u32) { + let pack_id = lane * packs_per_lane + _pi; + let w_row = pack_id / packs_per_lane; // 0..31 (BN rows) + let pack_col = pack_id % packs_per_lane; // which u32 in the BK=16 slice let pack_dev = w_expert_base + (n_tile_base + w_row) * packs_per_row - + kb / 8u32 + + kb / vals_per_pack + pack_col; let packed = load(w[pack_dev]); - let k_off = kb + pack_col * 8u32; + let k_off = kb + pack_col * vals_per_pack; let g = k_off / group_size; let sb_off = sb_expert_base + (n_tile_base + w_row) * groups_per_row + g; let s = load(scales[sb_off]).cast::(); let b = load(biases[sb_off]).cast::(); - let dst = w_row * 16u32 + pack_col * 8u32; - for _j in range(0u32, 8u32, 1u32) { - let q = ((packed >> (_j * 4u32)) & 15u32).cast::(); + let dst = w_row * 16u32 + pack_col * vals_per_pack; + for _j in range(0u32, vals_per_pack, 1u32) { + let q = ((packed >> (_j * BITS)) & mask).cast::(); threadgroup_store("ws", dst + _j, s * q + b); } } @@ -225,22 +224,33 @@ mod tests { assert_eq!(setup, DType::F16, "bf16 activation must stage as half for matmul2d"); } - /// Codegen sanity — the MPP header + descriptor land in the MSL. + /// Codegen sanity — the MPP header + descriptor land in the MSL, for both + /// the int4 and int8 variants. #[test] fn codegen_emits_mpp_include() { - let mut k = mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(DType::F32); - k.name = "mt_moe_gather_qmm_mma_int4_bm16_mpp_f32".into(); - let msl = MslGenerator::default().generate(&k).expect("codegen"); - assert!(msl.contains("MetalPerformancePrimitives/MetalPerformancePrimitives.h")); - assert!(msl.contains("mpp::tensor_ops::matmul2d_descriptor")); - assert!(msl.contains("kernel void mt_moe_gather_qmm_mma_int4_bm16_mpp_f32")); + for (mut k, name) in [ + ( + mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(DType::F32), + "mt_moe_gather_qmm_mma_int4_bm16_mpp_f32", + ), + ( + mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(DType::F32), + "mt_moe_gather_qmm_mma_int8_bm16_mpp_f32", + ), + ] { + k.name = name.into(); + let msl = MslGenerator::default().generate(&k).expect("codegen"); + assert!(msl.contains("MetalPerformancePrimitives/MetalPerformancePrimitives.h")); + assert!(msl.contains("mpp::tensor_ops::matmul2d_descriptor")); + assert!(msl.contains(&format!("kernel void {name}"))); + } } } -/// New-syntax correctness test for the MPP MoE int4 BGEMM (BM=16). Oracle is -/// the clean per-row-`indices` dequant-then-grouped-matmul: each row `t` -/// resolves its expert from `indices[t]`, dequantizes that expert's int4 -/// weight (8 nibbles/u32, per-group scale/bias), and dots against the row's +/// New-syntax correctness tests for the MPP MoE BGEMM (BM=16), int4 + int8. +/// Oracle is the clean per-row-`indices` dequant-then-grouped-matmul: each row +/// `t` resolves its expert from `indices[t]`, dequantizes that expert's weight +/// (`32/BITS` codes/u32, per-group scale/bias), and dots against the row's /// input. Inputs are dtype-rounded so the GPU sees exactly what the oracle /// computes; tolerance is wide because the MPP cooperative-tensor accumulator /// reorders the K reduction. @@ -249,12 +259,12 @@ mod tests { pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::mt_moe_gather_qmm_mma_int4_bm16_mpp; - use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup}; + use super::{mt_moe_gather_qmm_mma_int4_bm16_mpp, mt_moe_gather_qmm_mma_int8_bm16_mpp}; + use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup, int8_indexed_setup}; + // Clean tile: BM=16 → ceil(64/16)=4 m-tiles, BN=32 → 64/32=2 n-tiles. #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int4_bm16_mpp(dt: DType) -> TestSetup { - // Clean tile: BM=16 → ceil(64/16)=4 m-tiles, BN=32 → 64/32=2 n-tiles. int4_indexed_setup( mt_moe_gather_qmm_mma_int4_bm16_mpp::kernel_ir_for(dt), MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, @@ -264,13 +274,26 @@ pub mod kernel_tests { dt, ) } + + #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] + fn test_moe_gather_qmm_mma_int8_bm16_mpp(dt: DType) -> TestSetup { + int8_indexed_setup( + mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt), + MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, + 32, // bn + 16, // bm + 32, // tpg + dt, + ) + } } -/// New-syntax benchmark for the MPP MoE int4 BGEMM (BM=16). Qwen3.6-A3B-ish. +/// New-syntax benchmarks for the MPP MoE BGEMM (BM=16), int4 + int8. +/// Qwen3.6-A3B-ish. pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::mt_moe_gather_qmm_mma_int4_bm16_mpp; + use super::{mt_moe_gather_qmm_mma_int4_bm16_mpp, mt_moe_gather_qmm_mma_int8_bm16_mpp}; use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] @@ -291,4 +314,23 @@ pub mod kernel_benches { dt, ) } + + #[bench(dtypes = [f32, f16, bf16])] + fn bench_moe_gather_qmm_mma_int8_bm16_mpp(dt: DType) -> BenchSetup { + int4_mma_bench( + mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt), + MmaBenchShape { + bits: 8, + bn: 32, + bm: 16, + tpg: 32, + m_total: 1024, + n_out: 256, + k_in: 2048, + n_experts: 128, + group_size: 64, + }, + dt, + ) + } } diff --git a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs index 17d095f4..f3f7d557 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs @@ -35,11 +35,12 @@ use metaltile::kernel; -/// MPP MoE int4 grouped BGEMM, BM=BN=64 / BK=32, 4 simdgroups (2×2). -/// Signature matches `…_bm16_mpp`. -#[kernel] +/// MPP MoE grouped BGEMM, BM=BN=64 / BK=32, 4 simdgroups (2×2). `BITS` ∈ {4, 8} +/// → `mt_moe_gather_qmm_mma_int4_bm64_mpp` / `_int8_bm64_mpp`. Same weight-unpack +/// fold as `moe_mpp` (`32/BITS` codes/u32), here over a BK=32 slice. +#[kernel(variants(BITS = [4, 8], suffix = "int{BITS}_bm64_mpp"))] #[allow(clippy::too_many_arguments)] -pub fn mt_moe_gather_qmm_mma_int4_bm64_mpp( +pub fn mt_moe_gather_qmm_mma( x: Tensor, w: Tensor, scales: Tensor, @@ -58,7 +59,9 @@ pub fn mt_moe_gather_qmm_mma_int4_bm64_mpp( // 2×2 warp grid: sm/sn select this SG's 32×32 sub-tile. let sg_m_base = (sg / 2u32) * 32u32; let sg_n_base = (sg & 1u32) * 32u32; - let packs_per_row = k_in / 8u32; + // Weight packing: `32/BITS` codes per u32 (int4 → 8 nibbles, int8 → 4 bytes). + let vals_per_pack = 32u32 / BITS; + let packs_per_row = k_in / vals_per_pack; let groups_per_row = k_in / group_size; // X coop-load: 128 lanes × 16 contiguous K = 2048 = BM(64)×TG_LD(32). let x_m_row = lane_in_tg / 2u32; @@ -119,24 +122,28 @@ pub fn mt_moe_gather_qmm_mma_int4_bm64_mpp( let xv = load(x[x_dev_base + _i]).cast::(); threadgroup_store("Xs", x_ws_base + _i, select(in_run_x, xv, 0.0f32)); } - // Dequant W → Ws. 128 lanes × 2 packs/lane = 256 packs. - for _pi in range(0u32, 2u32, 1u32) { - let pack_id = lane_in_tg * 2u32 + _pi; - let w_row = pack_id / 4u32; // 0..63 (BN rows) - let pack_in_row = pack_id & 3u32; // 0..3 (BK=32 → 4 packs) + // Dequant W → Ws. 128 lanes × `packs_per_lane` packs/lane; + // `vals_per_pack` codes/pack over the BK=32 slice. + let packs_in_row = 32u32 / vals_per_pack; // BK=32 → packs per BN row + let packs_per_lane = 16u32 / vals_per_pack; // (BN*BK/vals)/128 lanes + let mask = (1u32 << BITS) - 1u32; + for _pi in range(0u32, packs_per_lane, 1u32) { + let pack_id = lane_in_tg * packs_per_lane + _pi; + let w_row = pack_id / packs_in_row; // 0..63 (BN rows) + let pack_in_row = pack_id % packs_in_row; // which u32 in the BK=32 slice let pack_dev = w_expert_base + (n_tile_base + w_row) * packs_per_row - + kb / 8u32 + + kb / vals_per_pack + pack_in_row; let packed = load(w[pack_dev]); - let k_off = kb + pack_in_row * 8u32; + let k_off = kb + pack_in_row * vals_per_pack; let g = k_off / group_size; let sb_off = sb_expert_base + (n_tile_base + w_row) * groups_per_row + g; let s = load(scales[sb_off]).cast::(); let b = load(biases[sb_off]).cast::(); - let ws_base = w_row * 32u32 + pack_in_row * 8u32; - for _j in range(0u32, 8u32, 1u32) { - let q = ((packed >> (_j * 4u32)) & 15u32).cast::(); + let ws_base = w_row * 32u32 + pack_in_row * vals_per_pack; + for _j in range(0u32, vals_per_pack, 1u32) { + let q = ((packed >> (_j * BITS)) & mask).cast::(); threadgroup_store("Ws", ws_base + _j, s * q + b); } } @@ -219,28 +226,34 @@ mod tests { pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::mt_moe_gather_qmm_mma_int4_bm64_mpp; - use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup}; + use super::{mt_moe_gather_qmm_mma_int4_bm64_mpp, mt_moe_gather_qmm_mma_int8_bm64_mpp}; + use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup, int8_indexed_setup}; + // BN=64 → 64/64=1 n-tile, BM=64 → ceil(64/64)=1 m-tile. #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int4_bm64_mpp(dt: DType) -> TestSetup { - // BN=64 → 64/64=1 n-tile, BM=64 → ceil(64/64)=1 m-tile. int4_indexed_setup( mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(dt), MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, - 64, // bn - 64, // bm - 128, // tpg (4 SGs) - dt, + 64, 64, 128, dt, + ) + } + + #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] + fn test_moe_gather_qmm_mma_int8_bm64_mpp(dt: DType) -> TestSetup { + int8_indexed_setup( + mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(dt), + MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, + 64, 64, 128, dt, ) } } -/// New-syntax benchmark for the MPP MoE int4 BGEMM (BM=BN=64). Qwen3.6-A3B-ish. +/// New-syntax benchmarks for the MPP MoE BGEMM (BM=BN=64), int4 + int8. pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::mt_moe_gather_qmm_mma_int4_bm64_mpp; + use super::{mt_moe_gather_qmm_mma_int4_bm64_mpp, mt_moe_gather_qmm_mma_int8_bm64_mpp}; use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] @@ -248,15 +261,20 @@ pub mod kernel_benches { int4_mma_bench( mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(dt), MmaBenchShape { - bits: 4, - bn: 64, - bm: 64, - tpg: 128, - m_total: 1024, - n_out: 256, - k_in: 2048, - n_experts: 128, - group_size: 64, + bits: 4, bn: 64, bm: 64, tpg: 128, + m_total: 1024, n_out: 256, k_in: 2048, n_experts: 128, group_size: 64, + }, + dt, + ) + } + + #[bench(dtypes = [f32, f16, bf16])] + fn bench_moe_gather_qmm_mma_int8_bm64_mpp(dt: DType) -> BenchSetup { + int4_mma_bench( + mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(dt), + MmaBenchShape { + bits: 8, bn: 64, bm: 64, tpg: 128, + m_total: 1024, n_out: 256, k_in: 2048, n_experts: 128, group_size: 64, }, dt, ) diff --git a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64_int8.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64_int8.rs deleted file mode 100644 index 7691f903..00000000 --- a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64_int8.rs +++ /dev/null @@ -1,293 +0,0 @@ -//! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric -//! SPDX-License-Identifier: Apache-2.0 -//! MPP-backed MoE grouped int8 BGEMM — `mt_moe_gather_qmm_mma_int8_bm64_mpp`. -//! -//! BM=BN=64, BK=32 int8 variant of `mt_moe_gather_qmm_mma_int4_bm64_mpp`. -//! Runs **4 simdgroups** in a 2×2 warp grid over a 64×64 tile — each SG owns -//! a 32×32 sub-tile and a 32×32×32 `matmul2d`. For long-context prefill the -//! larger tile amortises the int8 dequant across more output. -//! -//! ## int4 → int8 lane mapping (BM=64) -//! -//! W tile size: BN(64) × BK(32) = 2048 elements. -//! -//! - **int4**: 128 lanes × 2 packs/lane × 8 nibbles/pack = 2048 ✓ -//! - pack_id = lane_in_tg*2 + _pi; w_row = pack_id/4; pack_in_row = pack_id%4 -//! - k_off = kb + pack_in_row*8; ws_base = w_row*32 + pack_in_row*8 -//! - Extracts 8 nibbles: `(packed >> (j*4)) & 0xf` -//! -//! - **int8**: 128 lanes × 4 packs/lane × 4 bytes/pack = 2048 ✓ -//! - pack_id = lane_in_tg*4 + _pi; w_row = pack_id/8; pack_in_row = pack_id%8 -//! - k_off = kb + pack_in_row*4; ws_base = w_row*32 + pack_in_row*4 -//! - Extracts 4 bytes: `(packed >> (j*8)) & 0xff` -//! -//! ## Descriptor -//! -//! `matmul2d_descriptor(32, 32, 32, ta=false, tb=true, tc=false, -//! multiply_accumulate)` — all dims 32, so the inputs are cooperative -//! tensors (not the direct-input path the `…_bm8` variant needs). -//! -//! ## bf16 staging -//! -//! `coop_stage(T)` = `half` for `T = bf16`, else `T`. Apple's `matmul2d` -//! mishandles `bfloat` cooperative tensors; `half` losslessly covers -//! bf16's mantissa. Accumulation is fp32. -//! -//! ## Dispatch invariants -//! -//! - Mode `Reduction`; grid `[N/64, ceil(M/64), 1]`; threadgroup -//! `[128, 1, 1]` (4 simdgroups, 2×2 warp grid). -//! - `k_in % 32 == 0`, `n_out % 64 == 0`, `group_size` divides `k_in`. -//! -//! Correctness validated by `tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs`. - -use metaltile::kernel; - -/// MPP MoE int8 grouped BGEMM, BM=BN=64 / BK=32, 4 simdgroups (2×2). -/// Signature matches `…_int4_bm64_mpp`. -#[kernel] -#[allow(clippy::too_many_arguments)] -pub fn mt_moe_gather_qmm_mma_int8_bm64_mpp( - x: Tensor, - w: Tensor, - scales: Tensor, - biases: Tensor, - indices: Tensor, - mut out: Tensor, - #[constexpr] m_total: u32, - #[constexpr] n_out: u32, - #[constexpr] k_in: u32, - #[constexpr] group_size: u32, -) { - let n_tile_base = tgid_x * 64u32; - let m_tile_base = tgid_y * 64u32; - let sg = simd_group_id(); - let lane_in_tg = sg * 32u32 + simd_lane; - // 2×2 warp grid: sg_m_base / sg_n_base select this SG's 32×32 sub-tile. - let sg_m_base = (sg / 2u32) * 32u32; - let sg_n_base = (sg & 1u32) * 32u32; - // int8: 4 bytes per u32 → k_in / 4 packs per weight row. - let packs_per_row = k_in / 4u32; - let groups_per_row = k_in / group_size; - // X coop-load: 128 lanes × 16 contiguous K = 2048 = BM(64)×TG_LD(32). - let x_m_row = lane_in_tg / 2u32; - let x_k_base = (lane_in_tg & 1u32) * 16u32; - threadgroup_alloc("Xs", 2048, coop_stage(T)); // 64 × 32 - threadgroup_alloc("Ws", 2048, coop_stage(T)); // 64 × 32 - threadgroup_alloc("OutScratch", 4096, f32); // 4 SG × 32 × 32 - // Descriptor 32×32×32, cooperative-tensor inputs, accumulate. - coop_tile_setup( - "gemm", - 32, - 32, - 32, // m, n, k - coop_stage(T), - "accumulate", - "simdgroup", - f32, - false, - true, - false, - ); - let mut sub_offset = 0u32; - for _sub_iter in range(0u32, 64u32, 1u32) { - let cur_row = m_tile_base + sub_offset; - let cur_in_range = (sub_offset < 64u32) & (cur_row < m_total); - let cur_expert = select(cur_in_range, load(indices[cur_row]), 4294967295u32); - // Walk forward to find the first row whose expert differs, clamping - // sub_end at the tile boundary or at m_total. - let mut sub_end = 64u32; - let mut found = 0u32; - for _ii in range(0u32, 64u32, 1u32) { - let probe = sub_offset + 1u32 + _ii; - let probe_row = m_tile_base + probe; - let probe_in_range = (probe < 64u32) & (probe_row < m_total); - if probe_in_range & (found == 0u32) { - let e = load(indices[probe_row]); - if e != cur_expert { - sub_end = probe; - found = 1u32; - } - } - if (probe < 64u32) & (probe_row >= m_total) & (found == 0u32) { - sub_end = probe; - found = 1u32; - } - } - let cur_valid = (cur_expert != 4294967295u32) & (sub_offset < 64u32); - if cur_valid { - let w_expert_base = cur_expert * n_out * packs_per_row; - let sb_expert_base = cur_expert * n_out * groups_per_row; - coop_tile_zero("gemm"); - for kb in range(0u32, k_in, 32u32) { - // Stage X[m_tile_base..+64, kb..kb+32] → Xs. 128 lanes × 16. - let gr_x = m_tile_base + x_m_row; - let in_run_x = (x_m_row >= sub_offset) & (x_m_row < sub_end) & (gr_x < m_total); - let safe_gr_x = select(in_run_x, gr_x, 0u32); - let x_dev_base = safe_gr_x * k_in + kb + x_k_base; - let x_ws_base = x_m_row * 32u32 + x_k_base; - for _i in range(0u32, 16u32, 1u32) { - let xv = load(x[x_dev_base + _i]).cast::(); - threadgroup_store("Xs", x_ws_base + _i, select(in_run_x, xv, 0.0f32)); - } - // Dequant W → Ws. - // - // int8 lane mapping: 128 lanes × 4 packs/lane × 4 bytes/pack - // = 2048 = BN(64) × BK(32). - // - // pack_id = lane_in_tg*4 + _pi (0..511) - // w_row = pack_id / 8 (0..63 = BN rows) - // pack_in_row = pack_id % 8 (0..7 — BK=32 → 8 u32s of 4 bytes) - // - // k_off = kb + pack_in_row*4 - // ws_base = w_row*32 + pack_in_row*4 - // - // Each pack holds 4 bytes (one per K-element); inner _j in 0..4 - // extracts byte j via (packed >> (j*8)) & 0xff. - for _pi in range(0u32, 4u32, 1u32) { - let pack_id = lane_in_tg * 4u32 + _pi; - let w_row = pack_id / 8u32; // 0..63 (BN rows) - let pack_in_row = pack_id & 7u32; // 0..7 (BK=32 → 8 packs) - let pack_dev = w_expert_base - + (n_tile_base + w_row) * packs_per_row - + kb / 4u32 - + pack_in_row; - let packed = load(w[pack_dev]); - let k_off = kb + pack_in_row * 4u32; - let g = k_off / group_size; - let sb_off = sb_expert_base + (n_tile_base + w_row) * groups_per_row + g; - let s = load(scales[sb_off]).cast::(); - let b = load(biases[sb_off]).cast::(); - let ws_base = w_row * 32u32 + pack_in_row * 4u32; - for _j in range(0u32, 4u32, 1u32) { - let q = ((packed >> (_j * 8u32)) & 255u32).cast::(); - threadgroup_store("Ws", ws_base + _j, s * q + b); - } - } - threadgroup_barrier(); - // Per-SG 32×32 sub-tile views into Xs / Ws (offset by the - // SG's 32-row span × TG_LD=32). extents<32, 32> = K-inner. - coop_tile_load_a("gemm", "Xs", true, coop_stage(T), 32, 32, sg_m_base * 32u32); - coop_tile_load_b("gemm", "Ws", true, coop_stage(T), 32, 32, sg_n_base * 32u32); - coop_tile_run("gemm"); - threadgroup_barrier(); - } - // Store this SG's 32×32 fp32 result into its OutScratch slot. - coop_tile_store_c("gemm", "OutScratch", true, f32, 32, 32, sg * 1024u32); - threadgroup_barrier(); - // Coop-write OutScratch → out. 128 lanes × 32 = 4096 = BM*BN. - // Each (mr, nc) lives in SG `(mr/32)*2 + (nc/32)`'s scratch. - for _e in range(0u32, 32u32, 1u32) { - let flat = lane_in_tg * 32u32 + _e; - let mr = flat / 64u32; - let nc = flat & 63u32; - let gr = m_tile_base + mr; - let gc = n_tile_base + nc; - let in_run = (mr >= sub_offset) & (mr < sub_end) & (gr < m_total) & (gc < n_out); - if in_run { - let src_sg = (mr / 32u32) * 2u32 + nc / 32u32; - let v = threadgroup_load( - "OutScratch", - src_sg * 1024u32 + (mr & 31u32) * 32u32 + (nc & 31u32), - ); - store(out[gr * n_out + gc], v.cast::()); - } - } - threadgroup_barrier(); - } - sub_offset = sub_end; - } -} - -#[cfg(test)] -mod tests { - use metaltile::core::{DType, ir::Op}; - - use super::*; - - #[test] - fn kernel_ir_constructs_and_uses_coop_tile_ops() { - for dt in [DType::F32, DType::F16, DType::BF16] { - let k = mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(dt); - assert_eq!(k.params.len(), 6); - assert_eq!(k.constexprs.len(), 4); - let all_ops = - || std::iter::once(&k.body).chain(k.blocks.values()).flat_map(|b| b.ops.iter()); - assert!(!all_ops().any(|op| matches!(op, Op::InlineMsl { .. }))); - assert!(all_ops().any(|op| matches!(op, Op::CoopTileSetup { .. }))); - assert!(all_ops().any(|op| matches!(op, Op::CoopTileRun { .. }))); - } - } - - #[test] - fn bf16_stages_through_half() { - let k = mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(DType::BF16); - let setup = std::iter::once(&k.body) - .chain(k.blocks.values()) - .flat_map(|b| b.ops.iter()) - .find_map(|op| match op { - Op::CoopTileSetup { act_dtype, .. } => Some(*act_dtype), - _ => None, - }) - .expect("CoopTileSetup present"); - assert_eq!(setup, DType::F16, "bf16 activation must stage as half"); - } -} - -/// New-syntax correctness test for the MPP MoE int8 BGEMM (BM=BN=64, 4 SGs). -/// Oracle is the shared per-row-`indices` int8 dequant-then-grouped-matmul -/// (4 unsigned bytes per u32). Inputs are dtype-rounded; tolerance is wide -/// because the 2×2 warp-grid cooperative-tensor accumulator reorders the K -/// reduction. -/// -/// Grid (Reduction, 4 simdgroups per TG): `grid_3d(n_out/64, ceil(m_total/64), 1, [128,1,1])`. -pub mod kernel_tests { - use metaltile::{test::*, test_kernel}; - - use super::mt_moe_gather_qmm_mma_int8_bm64_mpp; - use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int8_indexed_setup}; - - #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] - fn test_moe_gather_qmm_mma_int8_bm64_mpp(dt: DType) -> TestSetup { - // BN=64 → 64/64=1 n-tile, BM=64 → ceil(64/64)=1 m-tile. BK=32 → k_in=64 - // is 2 K-blocks; group_size=32 aligns to the BK stride. - int8_indexed_setup( - mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(dt), - MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, - 64, // bn - 64, // bm - 128, // tpg (4 SGs) - dt, - ) - } -} - -/// New-syntax benchmark for the MPP MoE int8 BGEMM (BM=BN=64). `bits=8` → -/// `k_in/4` u32 weight words/row. -/// -/// Grid (Reduction, 4 simdgroups per TG): `grid_3d(n_out/64, ceil(m_total/64), 1, [128,1,1])`. -pub mod kernel_benches { - use metaltile::{bench, test::*}; - - use super::mt_moe_gather_qmm_mma_int8_bm64_mpp; - use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; - - #[bench(dtypes = [f32, f16, bf16])] - fn bench_moe_gather_qmm_mma_int8_bm64_mpp(dt: DType) -> BenchSetup { - int4_mma_bench( - mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(dt), - MmaBenchShape { - bits: 8, - bn: 64, - bm: 64, - tpg: 128, - m_total: 1024, - n_out: 256, - k_in: 2048, - n_experts: 128, - group_size: 64, - }, - dt, - ) - } -} diff --git a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs index ff6f8af4..63b81daf 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs @@ -35,11 +35,12 @@ use metaltile::kernel; -/// MPP MoE int4 grouped BGEMM, BM=8 / BN=32 / BK=16, one simdgroup, -/// direct-input `matmul2d`. Signature matches `…_bm16_mpp`. -#[kernel] +/// MPP MoE grouped BGEMM, BM=8 / BN=32 / BK=16, one simdgroup, direct-input +/// `matmul2d`. `BITS` ∈ {4, 8} → `mt_moe_gather_qmm_mma_int4_bm8_mpp` / +/// `_int8_bm8_mpp`. Same weight-unpack fold as `moe_mpp` (`32/BITS` codes/u32). +#[kernel(variants(BITS = [4, 8], suffix = "int{BITS}_bm8_mpp"))] #[allow(clippy::too_many_arguments)] -pub fn mt_moe_gather_qmm_mma_int4_bm8_mpp( +pub fn mt_moe_gather_qmm_mma( x: Tensor, w: Tensor, scales: Tensor, @@ -54,7 +55,9 @@ pub fn mt_moe_gather_qmm_mma_int4_bm8_mpp( let n_tile_base = tgid_x * 32u32; let m_tile_base = tgid_y * 8u32; let lane = simd_lane; - let packs_per_row = k_in / 8u32; + // Weight packing: `32/BITS` codes per u32 (int4 → 8 nibbles, int8 → 4 bytes). + let vals_per_pack = 32u32 / BITS; + let packs_per_row = k_in / vals_per_pack; let groups_per_row = k_in / group_size; threadgroup_alloc("xs", 128, coop_stage(T)); // 8 × 16 threadgroup_alloc("ws", 512, coop_stage(T)); // 32 × 16 @@ -121,24 +124,27 @@ pub fn mt_moe_gather_qmm_mma_int4_bm8_mpp( let xv = load(x[safe_g * k_in + kb + kc]).cast::(); threadgroup_store("xs", mr * 16u32 + kc, select(in_run, xv, 0.0f32)); } - // Dequant W → ws. 32 lanes × 2 packs/lane, 8 nibbles/pack. - for _pi in range(0u32, 2u32, 1u32) { - let pack_id = lane * 2u32 + _pi; - let w_row = pack_id / 2u32; - let pack_col = pack_id % 2u32; + // Dequant W → ws. 32 lanes × `packs_per_lane` packs/lane, + // `vals_per_pack` codes/pack. + let packs_per_lane = 16u32 / vals_per_pack; + let mask = (1u32 << BITS) - 1u32; + for _pi in range(0u32, packs_per_lane, 1u32) { + let pack_id = lane * packs_per_lane + _pi; + let w_row = pack_id / packs_per_lane; + let pack_col = pack_id % packs_per_lane; let pack_dev = w_expert_base + (n_tile_base + w_row) * packs_per_row - + kb / 8u32 + + kb / vals_per_pack + pack_col; let packed = load(w[pack_dev]); - let k_off = kb + pack_col * 8u32; + let k_off = kb + pack_col * vals_per_pack; let g = k_off / group_size; let sb_off = sb_expert_base + (n_tile_base + w_row) * groups_per_row + g; let s = load(scales[sb_off]).cast::(); let b = load(biases[sb_off]).cast::(); - let dst = w_row * 16u32 + pack_col * 8u32; - for _j in range(0u32, 8u32, 1u32) { - let q = ((packed >> (_j * 4u32)) & 15u32).cast::(); + let dst = w_row * 16u32 + pack_col * vals_per_pack; + for _j in range(0u32, vals_per_pack, 1u32) { + let q = ((packed >> (_j * BITS)) & mask).cast::(); threadgroup_store("ws", dst + _j, s * q + b); } } @@ -205,36 +211,42 @@ mod tests { } } -/// New-syntax correctness test for the MPP MoE int4 BGEMM (BM=8). Shares the -/// per-row-`indices` int4 dequant-then-matmul oracle with the BM=16 sibling; -/// only the m-tile height (BM=8 → ceil(M/8) m-tiles) differs. +/// New-syntax correctness tests for the MPP MoE BGEMM (BM=8), int4 + int8. +/// Shares the per-row-`indices` dequant-then-matmul oracle with the BM=16 +/// sibling; only the m-tile height (BM=8 → ceil(M/8) m-tiles) differs. /// /// Grid (Reduction, 1 simdgroup per TG): `grid_3d(n_out/32, ceil(m_total/8), 1, [32,1,1])`. pub mod kernel_tests { use metaltile::{test::*, test_kernel}; - use super::mt_moe_gather_qmm_mma_int4_bm8_mpp; - use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup}; + use super::{mt_moe_gather_qmm_mma_int4_bm8_mpp, mt_moe_gather_qmm_mma_int8_bm8_mpp}; + use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup, int8_indexed_setup}; + // BM=8 → ceil(64/8)=8 m-tiles, BN=32 → 64/32=2 n-tiles. #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] fn test_moe_gather_qmm_mma_int4_bm8_mpp(dt: DType) -> TestSetup { - // BM=8 → ceil(64/8)=8 m-tiles, BN=32 → 64/32=2 n-tiles. int4_indexed_setup( mt_moe_gather_qmm_mma_int4_bm8_mpp::kernel_ir_for(dt), MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, - 32, // bn - 8, // bm - 32, // tpg - dt, + 32, 8, 32, dt, + ) + } + + #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] + fn test_moe_gather_qmm_mma_int8_bm8_mpp(dt: DType) -> TestSetup { + int8_indexed_setup( + mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt), + MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, + 32, 8, 32, dt, ) } } -/// New-syntax benchmark for the MPP MoE int4 BGEMM (BM=8). Qwen3.6-A3B-ish. +/// New-syntax benchmarks for the MPP MoE BGEMM (BM=8), int4 + int8. pub mod kernel_benches { use metaltile::{bench, test::*}; - use super::mt_moe_gather_qmm_mma_int4_bm8_mpp; + use super::{mt_moe_gather_qmm_mma_int4_bm8_mpp, mt_moe_gather_qmm_mma_int8_bm8_mpp}; use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; #[bench(dtypes = [f32, f16, bf16])] @@ -242,15 +254,20 @@ pub mod kernel_benches { int4_mma_bench( mt_moe_gather_qmm_mma_int4_bm8_mpp::kernel_ir_for(dt), MmaBenchShape { - bits: 4, - bn: 32, - bm: 8, - tpg: 32, - m_total: 1024, - n_out: 256, - k_in: 2048, - n_experts: 128, - group_size: 64, + bits: 4, bn: 32, bm: 8, tpg: 32, + m_total: 1024, n_out: 256, k_in: 2048, n_experts: 128, group_size: 64, + }, + dt, + ) + } + + #[bench(dtypes = [f32, f16, bf16])] + fn bench_moe_gather_qmm_mma_int8_bm8_mpp(dt: DType) -> BenchSetup { + int4_mma_bench( + mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt), + MmaBenchShape { + bits: 8, bn: 32, bm: 8, tpg: 32, + m_total: 1024, n_out: 256, k_in: 2048, n_experts: 128, group_size: 64, }, dt, ) diff --git a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8_int8.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8_int8.rs deleted file mode 100644 index f5a10ad7..00000000 --- a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8_int8.rs +++ /dev/null @@ -1,288 +0,0 @@ -//! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric -//! SPDX-License-Identifier: Apache-2.0 -//! MPP-backed MoE grouped int8 BGEMM — `mt_moe_gather_qmm_mma_int8_bm8_mpp`. -//! -//! BM=8 int8 sibling of `mt_moe_gather_qmm_mma_int4_bm8_mpp`. Same algorithm -//! and call-site signature; the weight layout changes from int4 (8 nibbles/u32) -//! to int8 (4 bytes/u32), doubling the number of weight u32s per row but -//! halving the packing inner-loop work per word. -//! -//! ## Direct-input matmul2d -//! -//! Descriptor `matmul2d_descriptor(8, 32, 16, ta=false, tb=true, tc=false, -//! multiply_accumulate)`. With M=8 the inputs cannot be cooperative tensors -//! (Apple's MPP path requires at least one of M/N/K ≥ 16 for cooperative -//! tensor descriptors), so A and B are passed as **direct** `metal::tensor` -//! views over threadgroup memory — the `direct_inputs` form. -//! -//! ## int4 → int8 lane mapping (BM=8) -//! -//! W tile size: BN(32) × BK(16) = 512 elements. -//! -//! - **int4**: 32 lanes × 2 packs/lane × 8 nibbles/pack = 512 ✓ -//! - pack_id = lane*2 + _pi; w_row = pack_id/2; pack_col = pack_id%2 -//! - k_off = kb + pack_col*8; dst = w_row*16 + pack_col*8 -//! - Extracts 8 nibbles: `(packed >> (j*4)) & 0xf` -//! -//! - **int8**: 32 lanes × 4 packs/lane × 4 bytes/pack = 512 ✓ -//! - pack_id = lane*4 + _pi; w_row = pack_id/4; pack_col = pack_id%4 -//! - k_off = kb + pack_col*4; dst = w_row*16 + pack_col*4 -//! - Extracts 4 bytes: `(packed >> (j*8)) & 0xff` -//! -//! ## bf16 staging -//! -//! `coop_stage(T)` = `half` for `T = bf16`, else `T` — Apple's `matmul2d` -//! mishandles `bfloat` operands, and `half` losslessly covers bf16's -//! mantissa. Accumulation is fp32. -//! -//! ## Dispatch invariants -//! -//! - Mode `Reduction`; grid `[N/32, ceil(M/8), 1]`; threadgroup -//! `[32, 1, 1]` (1 simdgroup). -//! - `k_in % 16 == 0`, `n_out % 32 == 0`, `group_size` divides `k_in`. -//! -//! Correctness validated by `tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs`. - -use metaltile::kernel; - -/// MPP MoE int8 grouped BGEMM, BM=8 / BN=32 / BK=16, one simdgroup, -/// direct-input `matmul2d`. Signature matches `…_int4_bm8_mpp`. -#[kernel] -#[allow(clippy::too_many_arguments)] -pub fn mt_moe_gather_qmm_mma_int8_bm8_mpp( - x: Tensor, - w: Tensor, - scales: Tensor, - biases: Tensor, - indices: Tensor, - mut out: Tensor, - #[constexpr] m_total: u32, - #[constexpr] n_out: u32, - #[constexpr] k_in: u32, - #[constexpr] group_size: u32, -) { - let n_tile_base = tgid_x * 32u32; - let m_tile_base = tgid_y * 8u32; - let lane = simd_lane; - // int8: 4 bytes per u32 → k_in / 4 packs per weight row. - let packs_per_row = k_in / 4u32; - let groups_per_row = k_in / group_size; - threadgroup_alloc("xs", 128, coop_stage(T)); // 8 × 16 - threadgroup_alloc("ws", 512, coop_stage(T)); // 32 × 16 - threadgroup_alloc("out_scratch", 256, f32); // 8 × 32 - // Descriptor 8×32×16, direct-input (M=8 → not a cooperative tensor). - // direct_inputs=true; A view = [K=16, M=8], B view = [K=16, N=32]. - coop_tile_setup( - "gemm", - 8, - 32, - 16, // m, n, k - coop_stage(T), - "accumulate", - "simdgroup", - f32, - false, - true, - false, - true, // direct_inputs - true, - 16, - 8, // a: is_tg, ei, eo - true, - 16, - 32, // b: is_tg, ei, eo - ); - let mut sub_offset = 0u32; - for _sub_iter in range(0u32, 8u32, 1u32) { - let cur_row = m_tile_base + sub_offset; - let cur_in_range = (sub_offset < 8u32) & (cur_row < m_total); - let cur_expert = select(cur_in_range, load(indices[cur_row]), 4294967295u32); - // Walk forward to find the first row whose expert differs, clamping - // sub_end at the tile boundary or at m_total. - let mut sub_end = 8u32; - let mut found = 0u32; - for _ii in range(0u32, 8u32, 1u32) { - let probe = sub_offset + 1u32 + _ii; - let probe_row = m_tile_base + probe; - let probe_in_range = (probe < 8u32) & (probe_row < m_total); - if probe_in_range & (found == 0u32) { - let e = load(indices[probe_row]); - if e != cur_expert { - sub_end = probe; - found = 1u32; - } - } - if (probe < 8u32) & (probe_row >= m_total) & (found == 0u32) { - sub_end = probe; - found = 1u32; - } - } - let cur_valid = (cur_expert != 4294967295u32) & (sub_offset < 8u32); - if cur_valid { - let w_expert_base = cur_expert * n_out * packs_per_row; - let sb_expert_base = cur_expert * n_out * groups_per_row; - coop_tile_zero("gemm"); - for kb in range(0u32, k_in, 16u32) { - // Stage X[m_tile_base..+8, kb..kb+16] → xs. 32 lanes × 4. - for _e in range(0u32, 4u32, 1u32) { - let flat = lane * 4u32 + _e; - let mr = flat / 16u32; - let kc = flat % 16u32; - let gr = m_tile_base + mr; - let in_run = (mr >= sub_offset) & (mr < sub_end) & (gr < m_total); - let safe_g = select(in_run, gr, 0u32); - let xv = load(x[safe_g * k_in + kb + kc]).cast::(); - threadgroup_store("xs", mr * 16u32 + kc, select(in_run, xv, 0.0f32)); - } - // Dequant W → ws. - // - // int8 lane mapping: 32 lanes × 4 packs/lane × 4 bytes/pack - // = 512 = BN(32) × BK(16). - // - // pack_id = lane*4 + _pi (0..127 — covers 32 w_rows × 4 packs/row) - // w_row = pack_id / 4 (0..31 = BN rows) - // pack_col= pack_id % 4 (0..3 — selects which of the 4 u32s in BK) - // - // k_off = kb + pack_col*4 (byte-offset of this pack's first element) - // dst = w_row*16 + pack_col*4 (flat index into ws threadgroup buf) - // - // Each pack holds 4 bytes (one per K-element); inner _j in 0..4 - // extracts byte j via (packed >> (j*8)) & 0xff. - for _pi in range(0u32, 4u32, 1u32) { - let pack_id = lane * 4u32 + _pi; - let w_row = pack_id / 4u32; - let pack_col = pack_id % 4u32; - let pack_dev = w_expert_base - + (n_tile_base + w_row) * packs_per_row - + kb / 4u32 - + pack_col; - let packed = load(w[pack_dev]); - let k_off = kb + pack_col * 4u32; - let g = k_off / group_size; - let sb_off = sb_expert_base + (n_tile_base + w_row) * groups_per_row + g; - let s = load(scales[sb_off]).cast::(); - let b = load(biases[sb_off]).cast::(); - let dst = w_row * 16u32 + pack_col * 4u32; - for _j in range(0u32, 4u32, 1u32) { - let q = ((packed >> (_j * 8u32)) & 255u32).cast::(); - threadgroup_store("ws", dst + _j, s * q + b); - } - } - threadgroup_barrier(); - coop_tile_load_a("gemm", "xs", true, coop_stage(T), 16, 8, true); - coop_tile_load_b("gemm", "ws", true, coop_stage(T), 16, 32, true); - coop_tile_run("gemm", true); - threadgroup_barrier(); - } - // C [M=8, N=32] row-major → extents N,M = 32,8. - coop_tile_store_c("gemm", "out_scratch", true, f32, 32, 8); - threadgroup_barrier(); - // Coop-write out_scratch → out. 32 lanes × 8 elems = 256 = BM*BN. - for _e in range(0u32, 8u32, 1u32) { - let flat = lane * 8u32 + _e; - let mr = flat / 32u32; - let nc = flat % 32u32; - let gr = m_tile_base + mr; - let gc = n_tile_base + nc; - let in_run = (mr >= sub_offset) & (mr < sub_end) & (gr < m_total) & (gc < n_out); - if in_run { - let v = threadgroup_load("out_scratch", mr * 32u32 + nc); - store(out[gr * n_out + gc], v.cast::()); - } - } - threadgroup_barrier(); - } - sub_offset = sub_end; - } -} - -#[cfg(test)] -mod tests { - use metaltile::core::{DType, ir::Op}; - - use super::*; - - #[test] - fn kernel_ir_constructs_and_uses_coop_tile_ops() { - for dt in [DType::F32, DType::F16, DType::BF16] { - let k = mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt); - assert_eq!(k.params.len(), 6); - assert_eq!(k.constexprs.len(), 4); - let all_ops = - || std::iter::once(&k.body).chain(k.blocks.values()).flat_map(|b| b.ops.iter()); - assert!(!all_ops().any(|op| matches!(op, Op::InlineMsl { .. }))); - assert!(all_ops().any(|op| matches!(op, Op::CoopTileSetup { .. }))); - assert!(all_ops().any(|op| matches!(op, Op::CoopTileRun { .. }))); - } - } - - #[test] - fn bf16_stages_through_half() { - let k = mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(DType::BF16); - let setup = std::iter::once(&k.body) - .chain(k.blocks.values()) - .flat_map(|b| b.ops.iter()) - .find_map(|op| match op { - Op::CoopTileSetup { act_dtype, .. } => Some(*act_dtype), - _ => None, - }) - .expect("CoopTileSetup present"); - assert_eq!(setup, DType::F16, "bf16 activation must stage as half"); - } -} - -/// New-syntax correctness test for the MPP MoE int8 BGEMM (BM=8). Oracle is the -/// shared per-row-`indices` int8 dequant-then-grouped-matmul (4 unsigned bytes -/// per u32). Inputs are dtype-rounded; tolerance is wide because the MPP -/// cooperative-tensor accumulator reorders the K reduction. -/// -/// Grid (Reduction, 1 simdgroup per TG): `grid_3d(n_out/32, ceil(m_total/8), 1, [32,1,1])`. -pub mod kernel_tests { - use metaltile::{test::*, test_kernel}; - - use super::mt_moe_gather_qmm_mma_int8_bm8_mpp; - use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int8_indexed_setup}; - - #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] - fn test_moe_gather_qmm_mma_int8_bm8_mpp(dt: DType) -> TestSetup { - // BM=8 → ceil(64/8)=8 m-tiles, BN=32 → 64/32=2 n-tiles. - int8_indexed_setup( - mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt), - MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, - 32, // bn - 8, // bm - 32, // tpg - dt, - ) - } -} - -/// New-syntax benchmark for the MPP MoE int8 BGEMM (BM=8). `bits=8` → -/// `k_in/4` u32 weight words/row. -/// -/// Grid (Reduction, 1 simdgroup per TG): `grid_3d(n_out/32, ceil(m_total/8), 1, [32,1,1])`. -pub mod kernel_benches { - use metaltile::{bench, test::*}; - - use super::mt_moe_gather_qmm_mma_int8_bm8_mpp; - use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; - - #[bench(dtypes = [f32, f16, bf16])] - fn bench_moe_gather_qmm_mma_int8_bm8_mpp(dt: DType) -> BenchSetup { - int4_mma_bench( - mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt), - MmaBenchShape { - bits: 8, - bn: 32, - bm: 8, - tpg: 32, - m_total: 1024, - n_out: 256, - k_in: 2048, - n_experts: 128, - group_size: 64, - }, - dt, - ) - } -} diff --git a/crates/metaltile-std/src/kernels/moe/moe_mpp_int8.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_int8.rs deleted file mode 100644 index d4d2ba5d..00000000 --- a/crates/metaltile-std/src/kernels/moe/moe_mpp_int8.rs +++ /dev/null @@ -1,304 +0,0 @@ -//! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric -//! SPDX-License-Identifier: Apache-2.0 -//! MPP-backed MoE int8 grouped BGEMM — `mt_moe_gather_qmm_mma_int8_bm16_mpp`. -//! -//! Int8 analogue of `moe_mpp::mt_moe_gather_qmm_mma_int4_bm16_mpp`. Same -//! BM=16 / BN=32 / BK=16 MPP cooperative-tensor tiling; the only change is -//! the weight coop-dequant inner loop: -//! -//! int4: 32 lanes × 2 packs/lane × 8 nibbles/pack = 512 = BN×BK ✓ -//! int8: 32 lanes × 4 packs/lane × 4 bytes/pack = 512 = BN×BK ✓ -//! -//! Each uint32 holds 4 consecutive unsigned-byte codes in LSB-first order. -//! `packs_per_row = k_in / 4` (4 bytes/u32 vs 8 nibbles/u32 for int4). -//! -//! ## bf16 staging -//! -//! Same `coop_stage(T)` trick as the int4 MPP kernel: bf16 activations are -//! staged through `half` so `mpp::tensor_ops::matmul2d` sees a supported -//! cooperative-tensor dtype. Accumulation is fp32. -//! -//! ## Descriptor -//! -//! `matmul2d_descriptor(16, 32, 16, ta=false, tb=true, tc=false, -//! multiply_accumulate)` — identical to the int4 MPP descriptor; only the -//! threadgroup W tile contents differ. -//! -//! ## Dispatch invariants -//! -//! - Mode `Reduction`; grid `[N/32, ceil(M/16), 1]`; threadgroup -//! `[32, 1, 1]` (1 simdgroup — `matmul2d` is `execution_simdgroup`). -//! - `k_in % 16 == 0`, `n_out % 32 == 0`, `group_size` divides `k_in`, -//! `(k_in / 4) % 4 == 0` (i.e. `k_in % 16 == 0`). -//! - macOS 26+ / Metal 4; on older toolchains a zero-write stub is emitted. -//! -//! Correctness: `tests/moe_gather_qmm_mpp_int8_correctness.rs` (cosine ≥ 0.999). - -use metaltile::kernel; - -/// MPP MoE int8 grouped BGEMM, BM=16 / BN=32 / BK=16, one simdgroup. -/// -/// Params: `x [m_total, k_in]`, `w [n_experts, n_out, k_in/4]` (int8 -/// packed, 4 bytes/uint32), `scales`/`biases [n_experts, n_out, -/// k_in/group]`, `indices [m_total]` (per-row expert id), `out -/// [m_total, n_out]`. -#[kernel] -#[allow(clippy::too_many_arguments)] -pub fn mt_moe_gather_qmm_mma_int8_bm16_mpp( - x: Tensor, - w: Tensor, - scales: Tensor, - biases: Tensor, - indices: Tensor, - mut out: Tensor, - #[constexpr] m_total: u32, - #[constexpr] n_out: u32, - #[constexpr] k_in: u32, - #[constexpr] group_size: u32, -) { - let n_tile_base = tgid_x * 32u32; - let m_tile_base = tgid_y * 16u32; - let lane = simd_lane; - // int8: 4 bytes per u32 → packs_per_row = k_in / 4. - let packs_per_row = k_in / 4u32; - let groups_per_row = k_in / group_size; - // Threadgroup staging tiles. `coop_stage(T)` = half for bf16, else T. - // `out_scratch` is fp32: `coop_tile_store_c` destination must match the - // accumulator type. - threadgroup_alloc("xs", 256, coop_stage(T)); // 16 × 16 - threadgroup_alloc("ws", 512, coop_stage(T)); // 32 × 16 - threadgroup_alloc("out_scratch", 512, f32); // 16 × 32 - // MPP descriptor 16×32×16, ta=false tb=true tc=false, accumulate. - coop_tile_setup( - "gemm", - 16, - 32, - 16, // m, n, k - coop_stage(T), - "accumulate", - "simdgroup", - f32, - false, - true, - false, - ); - // Walk the BM=16 rows in contiguous-expert sub-runs (identical to int4 MPP). - let mut sub_offset = 0u32; - for _sub_iter in range(0u32, 16u32, 1u32) { - let cur_row = m_tile_base + sub_offset; - let cur_in_range = (sub_offset < 16u32) & (cur_row < m_total); - let cur_expert = select(cur_in_range, load(indices[cur_row]), 4294967295u32); - // Find run end — first row whose expert differs (or OOB). - let mut sub_end = 16u32; - let mut found = 0u32; - for _ii in range(0u32, 16u32, 1u32) { - let probe = sub_offset + 1u32 + _ii; - let probe_row = m_tile_base + probe; - let probe_in_range = (probe < 16u32) & (probe_row < m_total); - if probe_in_range & (found == 0u32) { - let e = load(indices[probe_row]); - if e != cur_expert { - sub_end = probe; - found = 1u32; - } - } - if (probe < 16u32) & (probe_row >= m_total) & (found == 0u32) { - sub_end = probe; - found = 1u32; - } - } - let cur_valid = (cur_expert != 4294967295u32) & (sub_offset < 16u32); - if cur_valid { - let w_expert_base = cur_expert * n_out * packs_per_row; - let sb_expert_base = cur_expert * n_out * groups_per_row; - coop_tile_zero("gemm"); - for kb in range(0u32, k_in, 16u32) { - // Stage X[m_tile_base..+16, kb..kb+16] → xs. 32 lanes × 8. - for _e in range(0u32, 8u32, 1u32) { - let flat = lane * 8u32 + _e; - let mr = flat / 16u32; - let kc = flat % 16u32; - let gr = m_tile_base + mr; - let in_run = (mr >= sub_offset) & (mr < sub_end) & (gr < m_total); - let safe_g = select(in_run, gr, 0u32); - let xv = load(x[safe_g * k_in + kb + kc]).cast::(); - threadgroup_store("xs", mr * 16u32 + kc, select(in_run, xv, 0.0f32)); - } - // Dequant W[expert, n_tile_base..+32, kb..kb+16] → ws. - // int8: 32 lanes × 4 packs/lane × 4 bytes/pack = 512 = BN×BK. - // Lane assignment: - // pack_id = lane * 4 + _pi (0..127, but we have 32 lanes so 4 iters) - // w_row = pack_id / 4 (0..31: which BN row) - // pack_col = pack_id % 4 (0..3: which uint32 in BK=16 slice) - // k_off = kb + pack_col * 4 (byte offset of first element in pack) - // - // 32 lanes × 4 iters × 4 bytes/pack = 512 elements = 32 rows × 16 cols ✓ - for _pi in range(0u32, 4u32, 1u32) { - let pack_id = lane * 4u32 + _pi; - let w_row = pack_id / 4u32; // 0..31 (BN rows) - let pack_col = pack_id % 4u32; // 0..3 (BK=16 → 4 packs × 4 bytes) - let pack_dev = w_expert_base - + (n_tile_base + w_row) * packs_per_row - + kb / 4u32 - + pack_col; - let packed = load(w[pack_dev]); - // k_off = byte offset of the first element in this pack within the row. - let k_off = kb + pack_col * 4u32; - let g = k_off / group_size; - let sb_off = sb_expert_base + (n_tile_base + w_row) * groups_per_row + g; - let s = load(scales[sb_off]).cast::(); - let b = load(biases[sb_off]).cast::(); - // Extract 4 unsigned byte codes (LSB-first). - let q0 = (packed & 255u32).cast::(); - let q1 = ((packed >> 8u32) & 255u32).cast::(); - let q2 = ((packed >> 16u32) & 255u32).cast::(); - let q3 = ((packed >> 24u32) & 255u32).cast::(); - // Write to threadgroup ws at the correct row/col position. - let dst = w_row * 16u32 + pack_col * 4u32; - threadgroup_store("ws", dst, s * q0 + b); - threadgroup_store("ws", dst + 1u32, s * q1 + b); - threadgroup_store("ws", dst + 2u32, s * q2 + b); - threadgroup_store("ws", dst + 3u32, s * q3 + b); - } - threadgroup_barrier(); - // A = xs [M=16, K=16] (ta=false → extents K,M = 16,16). - // B = ws [N=32, K=16] (tb=true → extents K,N = 16,32). - coop_tile_load_a("gemm", "xs", true, coop_stage(T), 16, 16); - coop_tile_load_b("gemm", "ws", true, coop_stage(T), 16, 32); - coop_tile_run("gemm"); - threadgroup_barrier(); - } - // C [M=16, N=32] row-major → extents N,M = 32,16. - coop_tile_store_c("gemm", "out_scratch", true, f32, 32, 16); - threadgroup_barrier(); - // Coop-write out_scratch → out with the per-row expert mask. - // 32 lanes × 16 elems = 512 = BM*BN. - for _e in range(0u32, 16u32, 1u32) { - let flat = lane * 16u32 + _e; - let mr = flat / 32u32; - let nc = flat % 32u32; - let gr = m_tile_base + mr; - let gc = n_tile_base + nc; - let in_run = (mr >= sub_offset) & (mr < sub_end) & (gr < m_total) & (gc < n_out); - if in_run { - let v = threadgroup_load("out_scratch", mr * 32u32 + nc); - store(out[gr * n_out + gc], v.cast::()); - } - } - threadgroup_barrier(); - } - sub_offset = sub_end; - } -} - -#[cfg(test)] -mod tests { - use metaltile::{ - codegen::msl::MslGenerator, - core::{DType, ir::Op}, - }; - - use super::*; - - #[test] - fn kernel_ir_constructs_and_uses_coop_tile_ops() { - for dt in [DType::F32, DType::F16, DType::BF16] { - let k = mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt); - assert_eq!(k.name, "mt_moe_gather_qmm_mma_int8_bm16_mpp"); - assert_eq!(k.params.len(), 6); - assert!(k.params[5].is_output); - assert_eq!(k.constexprs.len(), 4); - // No raw inline MSL — the matmul is CoopTile* ops. - let all_ops = - || std::iter::once(&k.body).chain(k.blocks.values()).flat_map(|b| b.ops.iter()); - assert!(!all_ops().any(|op| matches!(op, Op::InlineMsl { .. }))); - assert!(all_ops().any(|op| matches!(op, Op::CoopTileSetup { .. }))); - assert!(all_ops().any(|op| matches!(op, Op::CoopTileRun { .. }))); - } - } - - /// bf16 must stage through `half`: the `coop_stage(T)` tiles and - /// cooperative tensors resolve to `half`, never `bfloat`. - #[test] - fn bf16_stages_through_half() { - let k = mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(DType::BF16); - let setup = std::iter::once(&k.body) - .chain(k.blocks.values()) - .flat_map(|b| b.ops.iter()) - .find_map(|op| match op { - Op::CoopTileSetup { act_dtype, .. } => Some(*act_dtype), - _ => None, - }) - .expect("CoopTileSetup present"); - assert_eq!(setup, DType::F16, "bf16 activation must stage as half for matmul2d"); - } - - /// Codegen sanity — the MPP header + descriptor land in the MSL. - #[test] - fn codegen_emits_mpp_include() { - let mut k = mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(DType::F32); - k.name = "mt_moe_gather_qmm_mma_int8_bm16_mpp_f32".into(); - let msl = MslGenerator::default().generate(&k).expect("codegen"); - assert!(msl.contains("MetalPerformancePrimitives/MetalPerformancePrimitives.h")); - assert!(msl.contains("mpp::tensor_ops::matmul2d_descriptor")); - assert!(msl.contains("kernel void mt_moe_gather_qmm_mma_int8_bm16_mpp_f32")); - } -} - -/// New-syntax correctness test for the MPP MoE int8 BGEMM (BM=16). Oracle is -/// the per-row-`indices` int8 dequant-then-grouped-matmul (4 unsigned bytes per -/// u32, per-group scale/bias) shared with the bm8/bm64 int8 variants. Inputs -/// are dtype-rounded so the GPU sees exactly what the oracle computes; tolerance -/// is wide because the MPP cooperative-tensor accumulator reorders the K -/// reduction. -/// -/// Grid (Reduction, 1 simdgroup per TG): `grid_3d(n_out/32, ceil(m_total/16), 1, [32,1,1])`. -pub mod kernel_tests { - use metaltile::{test::*, test_kernel}; - - use super::mt_moe_gather_qmm_mma_int8_bm16_mpp; - use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int8_indexed_setup}; - - #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] - fn test_moe_gather_qmm_mma_int8_bm16_mpp(dt: DType) -> TestSetup { - // BM=16 → ceil(64/16)=4 m-tiles, BN=32 → 64/32=2 n-tiles. - int8_indexed_setup( - mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt), - MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, - 32, // bn - 16, // bm - 32, // tpg - dt, - ) - } -} - -/// New-syntax benchmark for the MPP MoE int8 BGEMM (BM=16). `bits=8` → -/// `k_in/4` u32 weight words/row. -/// -/// Grid (Reduction, 1 simdgroup per TG): `grid_3d(n_out/32, ceil(m_total/16), 1, [32,1,1])`. -pub mod kernel_benches { - use metaltile::{bench, test::*}; - - use super::mt_moe_gather_qmm_mma_int8_bm16_mpp; - use crate::kernels::moe::moe_mpp_shared::{MmaBenchShape, int4_mma_bench}; - - #[bench(dtypes = [f32, f16, bf16])] - fn bench_moe_gather_qmm_mma_int8_bm16_mpp(dt: DType) -> BenchSetup { - int4_mma_bench( - mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt), - MmaBenchShape { - bits: 8, - bn: 32, - bm: 16, - tpg: 32, - m_total: 1024, - n_out: 256, - k_in: 2048, - n_experts: 128, - group_size: 64, - }, - dt, - ) - } -} diff --git a/crates/metaltile-std/src/kernels/moe/moe_mpp_shared.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_shared.rs index 6e41b027..eb735688 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_mpp_shared.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_shared.rs @@ -2,14 +2,15 @@ //! SPDX-License-Identifier: Apache-2.0 //! Shared test/bench helpers for the MPP MoE grouped BGEMM family. //! -//! Every MPP MoE kernel (`moe_mpp{,_int8,_bm8,_bm8_int8,_bm64,_bm64_int8}`) -//! shares one ABI — `x, w, scales, biases, indices, out` plus the four +//! Every MPP MoE kernel (`moe_mpp{,_bm8,_bm64}`, each a `variants(BITS=[4,8])` +//! pair) shares one ABI — `x, w, scales, biases, indices, out` plus the four //! `{m_total, n_out, k_in, group_size}` constexprs — and the same math: //! per-row expert routing via `indices[t]`, dequant-then-grouped-matmul. -//! Only the tile geometry (BM/BN/BK, SG count) and the weight bit-width -//! differ. These helpers centralise the int4 dequant oracle, the -//! per-variant `TestSetup`, and the per-variant `BenchSetup` so each -//! kernel file stays a thin shape-binding wrapper. +//! Only the tile geometry (BM/BN/BK, SG count) differs per file; the int4/int8 +//! weight bit-width is folded onto the `BITS` axis within each file. These +//! helpers centralise the int4 and int8 dequant oracles, the per-variant +//! `TestSetup`, and the per-variant `BenchSetup` so each kernel file stays a +//! thin shape-binding wrapper. use metaltile::{ core::{DType, ir::Kernel}, diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs index 2a54370b..a246e82a 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `kernels::moe::moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp`. +//! GPU correctness for `kernels::moe::moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp`. //! //! BM=BN=64 MPP MoE int8 kernel — same output semantics as the int4 BM=64 //! sibling but the weight layout changes from 8 nibbles/u32 to 4 bytes/u32. @@ -21,7 +21,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{moe_gather_qmm::mt_moe_gather_qmm_b8, moe_mpp_bm64_int8}; +use metaltile_std::kernels::moe::{moe_gather_qmm::mt_moe_gather_qmm_b8, moe_mpp_bm64}; /// Pack a row of int8 weight codes into uint32s (4 codes per uint, LE byte /// order). Code values must be in 0..=255. @@ -134,7 +134,7 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_matches_b8_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( + let mut k = moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( Dt::F32.to_dtype(), ); k.mode = KernelMode::Reduction; @@ -250,7 +250,7 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_matches_b8_multi_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( + let mut k = moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( Dt::F32.to_dtype(), ); k.mode = KernelMode::Reduction; @@ -370,7 +370,7 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_bf16_matches_b8_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm64_int8::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( + let mut k = moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( Dt::Bf16.to_dtype(), ); k.mode = KernelMode::Reduction; diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs index 2b1ff04c..fd9b42e2 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `kernels::moe::moe_mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp`. +//! GPU correctness for `kernels::moe::moe_mpp_bm8::mt_moe_gather_qmm_mma_int8_bm8_mpp`. //! //! BM=8 MPP MoE int8 kernel — same output semantics as the int4 BM=8 sibling //! but the weight layout changes from 8 nibbles/u32 to 4 bytes/u32. @@ -22,7 +22,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::{moe_gather_qmm::mt_moe_gather_qmm_b8, moe_mpp_bm8_int8}; +use metaltile_std::kernels::moe::{moe_gather_qmm::mt_moe_gather_qmm_b8, moe_mpp_bm8}; /// Pack a row of int8 weight codes into uint32s (4 codes per uint, LE byte /// order). Code values must be in 0..=255. @@ -174,7 +174,7 @@ fn run_case(case: &Case) { buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); let mut k = - moe_mpp_bm8_int8::mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt.to_dtype()); + moe_mpp_bm8::mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [ceil(N/32), ceil(T/8), 1]. TG: 32 lanes = 1 SG. let r = ctx diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs index 45ff3364..19c2433e 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_int8_correctness.rs @@ -2,7 +2,7 @@ //! SPDX-License-Identifier: Apache-2.0 #![allow(clippy::manual_is_multiple_of)] -//! GPU correctness for `kernels::moe::moe_mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp`. +//! GPU correctness for `kernels::moe::moe_mpp::mt_moe_gather_qmm_mma_int8_bm16_mpp`. //! //! MPP (MetalPerformancePrimitives) int8 MoE BGEMM — same algorithm as //! `mt_moe_gather_qmm_mma_int4_bm16_mpp` but with pack-aligned 8-bit @@ -21,7 +21,7 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::moe_mpp_int8; +use metaltile_std::kernels::moe::moe_mpp; // ── helpers ──────────────────────────────────────────────────────────────── @@ -95,7 +95,7 @@ fn skip_unless_apple10() -> bool { let family = probe.chip_family(); if family.is_none_or(|lvl| lvl < 10) { eprintln!( - "skip moe_mpp_int8: needs Apple10+ GPU (chip_family={family:?}); \ + "skip moe_mpp: needs Apple10+ GPU (chip_family={family:?}); \ kernel emits a zero-write stub on older silicon" ); true @@ -132,7 +132,7 @@ fn run_mpp_int8( buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("Context::new"); - let mut k = moe_mpp_int8::mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt.to_dtype()); + let mut k = moe_mpp::mt_moe_gather_qmm_mma_int8_bm16_mpp::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [N/BN=32, ceil(T/BM=16), 1], TG: [32, 1, 1] (1 SG — MPP matmul2d). let r = ctx From abcdf0fc7dde0b5a616460a04bb930c8b7c360b5 Mon Sep 17 00:00:00 2001 From: Eric Kryski <599019+ekryski@users.noreply.github.com> Date: Sat, 13 Jun 2026 14:21:35 -0600 Subject: [PATCH 5/7] docs: note moe mpp BITS-variant fold in the plan --- docs/specs/KERNEL_CONSOLIDATION_PLAN.md | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md index 4e210cbd..a67106e8 100644 --- a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md +++ b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md @@ -56,8 +56,9 @@ crates/metaltile-std/src/kernels/ 2pass/batched/sink) · multi(+d256/tree-mask) · prefill_mma · flash_quantized · aura_flash · steel/attn moe/ ✅ DONE — router_topk · permute (+unpermute) · gather_qmm (per-expert BGEMM) · - router_topk_biased / sigmoid_bias / sqrtsoftplus · mpp(bm8/bm64 × int8 × - block_scaled) + mpp_shared · bgemm/gemv (q2k/iq2xxs/q4, view/ws/rows) · gather_q4 · + router_topk_biased / sigmoid_bias / sqrtsoftplus · mpp{,_bm8,_bm64} (int4+int8 + folded onto a BITS variant axis) + mpp_*_block_scaled + mpp_shared · + bgemm/gemv (q2k/iq2xxs/q4, view/ws/rows) · gather_q4 · down_swiglu_accum / down_weighted_sum · dequant_gemv_expert_indexed(_block_scaled) · block_scaled. Filenames keep the moe_ prefix (match kernel names); format-axis fold deferred (§7). (orchestration split into router_topk / permute / gather_qmm.) From a76a11352ff2744b7e85eba95fe50303ad088652 Mon Sep 17 00:00:00 2001 From: Eric Kryski <599019+ekryski@users.noreply.github.com> Date: Sat, 13 Jun 2026 22:56:31 -0600 Subject: [PATCH 6/7] style: cargo fmt --all --- crates/metaltile-std/src/kernels/moe/mod.rs | 28 +++++++------- .../src/kernels/moe/moe_gather_qmm.rs | 1 - .../metaltile-std/src/kernels/moe/moe_mpp.rs | 6 ++- .../src/kernels/moe/moe_mpp_bm64.rs | 38 +++++++++++++++---- .../src/kernels/moe/moe_mpp_bm8.rs | 38 +++++++++++++++---- .../src/kernels/moe/moe_permute.rs | 1 - .../kernels/moe/moe_router_sigmoid_bias.rs | 6 +-- .../src/kernels/moe/moe_router_topk.rs | 3 -- .../tests/dsv4_router_topk_correctness.rs | 5 ++- ...oe_gather_qmm_mpp_bm64_int8_correctness.rs | 15 +++----- ...moe_gather_qmm_mpp_bm8_int8_correctness.rs | 3 +- 11 files changed, 93 insertions(+), 51 deletions(-) diff --git a/crates/metaltile-std/src/kernels/moe/mod.rs b/crates/metaltile-std/src/kernels/moe/mod.rs index 03832afa..8976416a 100644 --- a/crates/metaltile-std/src/kernels/moe/mod.rs +++ b/crates/metaltile-std/src/kernels/moe/mod.rs @@ -14,41 +14,41 @@ //! format-axis fold (plan §7) is deferred. // Routing — top-k expert selection, permute/unpermute, router pre-scores. -pub mod moe_router_topk; -pub mod moe_permute; pub mod moe_gather_qmm; -pub mod moe_router_topk_biased; +pub mod moe_permute; pub mod moe_router_sigmoid_bias; pub mod moe_router_sqrtsoftplus; +pub mod moe_router_topk; +pub mod moe_router_topk_biased; pub mod moe_sigmoid_bias; // MPP grouped BGEMM (one ABI, tile-geometry / bit-width variants; shared // test/bench helpers in `moe_mpp_shared`). pub mod moe_mpp; -pub mod moe_mpp_shared; -pub mod moe_mpp_bm8; -pub mod moe_mpp_bm64; pub mod moe_mpp_block_scaled; -pub mod moe_mpp_bm8_block_scaled; +pub mod moe_mpp_bm64; pub mod moe_mpp_bm64_block_scaled; +pub mod moe_mpp_bm8; +pub mod moe_mpp_bm8_block_scaled; +pub mod moe_mpp_shared; // GGUF-format per-expert matmul / matvec (q2k, iq2xxs, q4). +pub mod moe_bgemm_iq2xxs_bm64; +pub mod moe_bgemm_iq2xxs_mpp; +pub mod moe_bgemm_iq2xxs_view; +pub mod moe_bgemm_iq2xxs_view_u16_bm64; pub mod moe_bgemm_q2k_bm64; pub mod moe_bgemm_q2k_mpp; pub mod moe_bgemm_q2k_view; pub mod moe_bgemm_q2k_view_u16_bm64; pub mod moe_bgemm_q4_bm64; -pub mod moe_bgemm_iq2xxs_bm64; -pub mod moe_bgemm_iq2xxs_mpp; -pub mod moe_bgemm_iq2xxs_view; -pub mod moe_bgemm_iq2xxs_view_u16_bm64; pub mod moe_gather_down_q2k; pub mod moe_gather_gemv_iq2xxs; -pub mod moe_gemv_rows_q2k; pub mod moe_gemv_rows_iq2xxs; +pub mod moe_gemv_rows_q2k; pub mod moe_gemv_rows_view_iq2xxs; -pub mod moe_gemv_ws_q2k; pub mod moe_gemv_ws_iq2xxs; +pub mod moe_gemv_ws_q2k; // Batched Q4 expert gather (up / down / weighted-sum), seeded from gemv_q8. pub mod moe_gather_q4; @@ -58,6 +58,6 @@ pub mod moe_down_swiglu_accum; pub mod moe_down_weighted_sum_f16; // Expert-indexed dequant GEMV + block-scaled MoE matmul. +pub mod block_scaled_moe; pub mod dequant_gemv_expert_indexed; pub mod dequant_gemv_expert_indexed_block_scaled; -pub mod block_scaled_moe; diff --git a/crates/metaltile-std/src/kernels/moe/moe_gather_qmm.rs b/crates/metaltile-std/src/kernels/moe/moe_gather_qmm.rs index 8d723ab3..c2768619 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_gather_qmm.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_gather_qmm.rs @@ -3809,5 +3809,4 @@ pub mod kernel_benches { dt, ) } - } diff --git a/crates/metaltile-std/src/kernels/moe/moe_mpp.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp.rs index 527a12ac..36a0005b 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_mpp.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp.rs @@ -260,7 +260,11 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::{mt_moe_gather_qmm_mma_int4_bm16_mpp, mt_moe_gather_qmm_mma_int8_bm16_mpp}; - use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup, int8_indexed_setup}; + use crate::kernels::moe::moe_mpp_shared::{ + MmaTestShape, + int4_indexed_setup, + int8_indexed_setup, + }; // Clean tile: BM=16 → ceil(64/16)=4 m-tiles, BN=32 → 64/32=2 n-tiles. #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] diff --git a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs index f3f7d557..2613a03a 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm64.rs @@ -227,7 +227,11 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::{mt_moe_gather_qmm_mma_int4_bm64_mpp, mt_moe_gather_qmm_mma_int8_bm64_mpp}; - use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup, int8_indexed_setup}; + use crate::kernels::moe::moe_mpp_shared::{ + MmaTestShape, + int4_indexed_setup, + int8_indexed_setup, + }; // BN=64 → 64/64=1 n-tile, BM=64 → ceil(64/64)=1 m-tile. #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] @@ -235,7 +239,10 @@ pub mod kernel_tests { int4_indexed_setup( mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(dt), MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, - 64, 64, 128, dt, + 64, + 64, + 128, + dt, ) } @@ -244,7 +251,10 @@ pub mod kernel_tests { int8_indexed_setup( mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(dt), MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, - 64, 64, 128, dt, + 64, + 64, + 128, + dt, ) } } @@ -261,8 +271,15 @@ pub mod kernel_benches { int4_mma_bench( mt_moe_gather_qmm_mma_int4_bm64_mpp::kernel_ir_for(dt), MmaBenchShape { - bits: 4, bn: 64, bm: 64, tpg: 128, - m_total: 1024, n_out: 256, k_in: 2048, n_experts: 128, group_size: 64, + bits: 4, + bn: 64, + bm: 64, + tpg: 128, + m_total: 1024, + n_out: 256, + k_in: 2048, + n_experts: 128, + group_size: 64, }, dt, ) @@ -273,8 +290,15 @@ pub mod kernel_benches { int4_mma_bench( mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(dt), MmaBenchShape { - bits: 8, bn: 64, bm: 64, tpg: 128, - m_total: 1024, n_out: 256, k_in: 2048, n_experts: 128, group_size: 64, + bits: 8, + bn: 64, + bm: 64, + tpg: 128, + m_total: 1024, + n_out: 256, + k_in: 2048, + n_experts: 128, + group_size: 64, }, dt, ) diff --git a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs index 63b81daf..d7b6a029 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_mpp_bm8.rs @@ -220,7 +220,11 @@ pub mod kernel_tests { use metaltile::{test::*, test_kernel}; use super::{mt_moe_gather_qmm_mma_int4_bm8_mpp, mt_moe_gather_qmm_mma_int8_bm8_mpp}; - use crate::kernels::moe::moe_mpp_shared::{MmaTestShape, int4_indexed_setup, int8_indexed_setup}; + use crate::kernels::moe::moe_mpp_shared::{ + MmaTestShape, + int4_indexed_setup, + int8_indexed_setup, + }; // BM=8 → ceil(64/8)=8 m-tiles, BN=32 → 64/32=2 n-tiles. #[test_kernel(dtypes = [f32, f16, bf16], tol = [5e-3, 5e-2, 2e-1])] @@ -228,7 +232,10 @@ pub mod kernel_tests { int4_indexed_setup( mt_moe_gather_qmm_mma_int4_bm8_mpp::kernel_ir_for(dt), MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, - 32, 8, 32, dt, + 32, + 8, + 32, + dt, ) } @@ -237,7 +244,10 @@ pub mod kernel_tests { int8_indexed_setup( mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt), MmaTestShape { n_experts: 4, m_total: 64, n_out: 64, k_in: 64, group_size: 32 }, - 32, 8, 32, dt, + 32, + 8, + 32, + dt, ) } } @@ -254,8 +264,15 @@ pub mod kernel_benches { int4_mma_bench( mt_moe_gather_qmm_mma_int4_bm8_mpp::kernel_ir_for(dt), MmaBenchShape { - bits: 4, bn: 32, bm: 8, tpg: 32, - m_total: 1024, n_out: 256, k_in: 2048, n_experts: 128, group_size: 64, + bits: 4, + bn: 32, + bm: 8, + tpg: 32, + m_total: 1024, + n_out: 256, + k_in: 2048, + n_experts: 128, + group_size: 64, }, dt, ) @@ -266,8 +283,15 @@ pub mod kernel_benches { int4_mma_bench( mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt), MmaBenchShape { - bits: 8, bn: 32, bm: 8, tpg: 32, - m_total: 1024, n_out: 256, k_in: 2048, n_experts: 128, group_size: 64, + bits: 8, + bn: 32, + bm: 8, + tpg: 32, + m_total: 1024, + n_out: 256, + k_in: 2048, + n_experts: 128, + group_size: 64, }, dt, ) diff --git a/crates/metaltile-std/src/kernels/moe/moe_permute.rs b/crates/metaltile-std/src/kernels/moe/moe_permute.rs index 3597f445..8e8af2a4 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_permute.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_permute.rs @@ -117,7 +117,6 @@ pub fn mt_moe_permute( } } - pub mod kernel_tests { use metaltile::{test::*, test_kernel}; diff --git a/crates/metaltile-std/src/kernels/moe/moe_router_sigmoid_bias.rs b/crates/metaltile-std/src/kernels/moe/moe_router_sigmoid_bias.rs index a35f901d..7e55d21a 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_router_sigmoid_bias.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_router_sigmoid_bias.rs @@ -44,11 +44,7 @@ use metaltile::kernel; // a no-`` signature. The new declarative `#[bench]` on // `kernel_benches::bench_router` below handles registration directly. #[kernel] -pub fn mt_moe_router_sigmoid_bias( - logits: Tensor, - bias: Tensor, - mut scores: Tensor, -) { +pub fn mt_moe_router_sigmoid_bias(logits: Tensor, bias: Tensor, mut scores: Tensor) { let idx = tid; let l = load(logits[idx]); let b = load(bias[idx]); diff --git a/crates/metaltile-std/src/kernels/moe/moe_router_topk.rs b/crates/metaltile-std/src/kernels/moe/moe_router_topk.rs index d0725252..2bc01748 100644 --- a/crates/metaltile-std/src/kernels/moe/moe_router_topk.rs +++ b/crates/metaltile-std/src/kernels/moe/moe_router_topk.rs @@ -163,7 +163,6 @@ pub fn mt_moe_router_topk( } } - pub mod kernel_tests { use metaltile::{test::*, test_kernel}; @@ -237,7 +236,6 @@ pub mod kernel_tests { // norm_topk_prob = 0: raw global-softmax probs (Qwen3-Next). #[test_kernel(dtypes = [f32, f16, bf16], tol = [1e-4, 1e-2, 5e-2])] fn test_moe_router_topk_global(dt: DType) -> TestSetup { router_setup(dt, false) } - } pub mod kernel_benches { @@ -270,5 +268,4 @@ pub mod kernel_benches { .grid_3d(n_rows as u32, 1, 1, [32, 1, 1]) .bytes_moved(bytes as u64) } - } diff --git a/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs b/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs index a133ea47..4628c8b7 100644 --- a/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs +++ b/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs @@ -10,7 +10,10 @@ use std::collections::BTreeMap; use common::{Dt, gpu_lock, pack_bytes, unpack_bytes, unpack_u32_bytes}; use metaltile::{Context, core::ir::KernelMode}; -use metaltile_std::kernels::moe::moe_router_topk_biased::{mt_moe_router_topk_biased, mt_remap_u32}; +use metaltile_std::kernels::moe::moe_router_topk_biased::{ + mt_moe_router_topk_biased, + mt_remap_u32, +}; #[test] fn dsv4_router_topk_f32() { diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs index a246e82a..b8b44bda 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm64_int8_correctness.rs @@ -134,9 +134,8 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_matches_b8_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( - Dt::F32.to_dtype(), - ); + let mut k = + moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [ceil(N/64), ceil(T/64), 1]. TG: 128 lanes = 4 SGs (WM=WN=2). let r = ctx @@ -250,9 +249,8 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_matches_b8_multi_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( - Dt::F32.to_dtype(), - ); + let mut k = + moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(Dt::F32.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid( @@ -370,9 +368,8 @@ fn moe_gather_qmm_mma_int8_bm64_mpp_bf16_matches_b8_clean_tile() { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for( - Dt::Bf16.to_dtype(), - ); + let mut k = + moe_mpp_bm64::mt_moe_gather_qmm_mma_int8_bm64_mpp::kernel_ir_for(Dt::Bf16.to_dtype()); k.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid( diff --git a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs index fd9b42e2..363ac3df 100644 --- a/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs +++ b/crates/metaltile-std/tests/moe_gather_qmm_mpp_bm8_int8_correctness.rs @@ -173,8 +173,7 @@ fn run_case(case: &Case) { buffers.insert("k_in".into(), (k_in as u32).to_le_bytes().to_vec()); buffers.insert("group_size".into(), (group_size as u32).to_le_bytes().to_vec()); let ctx = Context::new().unwrap(); - let mut k = - moe_mpp_bm8::mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt.to_dtype()); + let mut k = moe_mpp_bm8::mt_moe_gather_qmm_mma_int8_bm8_mpp::kernel_ir_for(dt.to_dtype()); k.mode = KernelMode::Reduction; // Grid: [ceil(N/32), ceil(T/8), 1]. TG: 32 lanes = 1 SG. let r = ctx From 392b5a0d6ffe496f605fb1408d293e0b381f40c3 Mon Sep 17 00:00:00 2001 From: TheTom Date: Mon, 22 Jun 2026 16:28:50 -0500 Subject: [PATCH 7/7] test(moe): rename stale dsv4_router_topk test to moe_router_topk_biased The kernel was renamed mt_dsv4_router_topk -> mt_moe_router_topk_biased in this PR; the test file, test-fn name, and a stale ffai:: doc path still carried the old name. Rename-only, no assertion changes. --- ...k_correctness.rs => moe_router_topk_biased_correctness.rs} | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) rename crates/metaltile-std/tests/{dsv4_router_topk_correctness.rs => moe_router_topk_biased_correctness.rs} (96%) diff --git a/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs b/crates/metaltile-std/tests/moe_router_topk_biased_correctness.rs similarity index 96% rename from crates/metaltile-std/tests/dsv4_router_topk_correctness.rs rename to crates/metaltile-std/tests/moe_router_topk_biased_correctness.rs index 4628c8b7..b59e15d7 100644 --- a/crates/metaltile-std/tests/dsv4_router_topk_correctness.rs +++ b/crates/metaltile-std/tests/moe_router_topk_biased_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::mt_moe_router_topk_biased` — top-K by biased +//! GPU correctness for `kernels::moe::mt_moe_router_topk_biased` — top-K by biased //! score, weights = unbiased[chosen] renormalised to sum 1. #![cfg(target_os = "macos")] @@ -16,7 +16,7 @@ use metaltile_std::kernels::moe::moe_router_topk_biased::{ }; #[test] -fn dsv4_router_topk_f32() { +fn moe_router_topk_biased_f32() { let _g = gpu_lock(); let n = 256usize; let k = 6usize;