diff --git a/crates/metaltile-std/src/ffai/mod.rs b/crates/metaltile-std/src/ffai/mod.rs index 89da6dde..7d1c8ea6 100644 --- a/crates/metaltile-std/src/ffai/mod.rs +++ b/crates/metaltile-std/src/ffai/mod.rs @@ -44,8 +44,7 @@ pub mod flash_block_scaled_sdpa; pub mod flash_quantized_sdpa; pub mod gate_up_swiglu_fused; pub mod gelu_erf; -// gemm_q4_mpp / gemm_q8 / gemm_q8_mpp migrated to kernels/gemm/. -pub mod gemv_q8; +// gemm_q4_mpp / gemm_q8 / gemm_q8_mpp + gemv_q8 (split by family) → kernels/. pub mod gguf_dequant_iq2_xxs; pub mod gguf_dequant_iq2_xxs_raw; pub mod gguf_dequant_q2_k; diff --git a/crates/metaltile-std/src/kernels/convolution/conv1d_causal.rs b/crates/metaltile-std/src/kernels/convolution/conv1d_causal.rs index cc0dbcac..bc883040 100644 --- a/crates/metaltile-std/src/kernels/convolution/conv1d_causal.rs +++ b/crates/metaltile-std/src/kernels/convolution/conv1d_causal.rs @@ -96,6 +96,29 @@ pub fn mt_conv1d_causal_prefill( store(y[idx], acc * sig); } +// Causal-conv state roll (prefill->decode handoff for the short conv). +/// Roll a causal-conv state ON-DEVICE: `new = [old[conv_dim..], xbc]` (drop the +/// oldest conv_dim, append the current input) — keeps the Mamba conv history on +/// the GPU. `keep = (kc-2)*conv_dim`; indices clamped so both select branches +/// are in-bounds. +#[kernel] +pub fn mt_conv_roll( + old: Tensor, + xbc: Tensor, + mut newst: Tensor, + #[constexpr] conv_dim: u32, + #[constexpr] keep: u32, + #[constexpr] n: u32, +) { + let i = program_id::<0>(); + if i < n { + let oi = select(i < keep, i + conv_dim, 0u32); + let xi = select(i < keep, 0u32, i - keep); + let v = select(i < keep, load(old[oi]), load(xbc[xi])); + store(newst[i], v); + } +} + pub mod kernel_tests { use metaltile::{test::*, test_kernel}; diff --git a/crates/metaltile-std/src/ffai/gemv_q8.rs b/crates/metaltile-std/src/kernels/gemm/gemv_quantized.rs similarity index 65% rename from crates/metaltile-std/src/ffai/gemv_q8.rs rename to crates/metaltile-std/src/kernels/gemm/gemv_quantized.rs index a0da64f2..00cb9b60 100644 --- a/crates/metaltile-std/src/ffai/gemv_q8.rs +++ b/crates/metaltile-std/src/kernels/gemm/gemv_quantized.rs @@ -1,12 +1,18 @@ //! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric //! SPDX-License-Identifier: Apache-2.0 -//! Q8_0 inline-dequant GEMV — `out[r] = Σ_k dequant(W[r,k]) · x[k]`, -//! reading the Q8_0 weight straight from its split buffers so the dense -//! attention/shared-expert projections stay 1 byte/weight resident -//! instead of being pre-expanded to f16 (2 bytes). The DSv4 attention -//! block is bandwidth-bound on these projections (q_b / output_a / -//! output_b are Q8_0 on disk, ~100M weights/layer); halving their bytes -//! roughly halves the attn GPU time. +//! Inline-dequant GEMV — `out[r] = Σ_k dequant(W[r,k]) · x[k]` — for Q8_0 (1 +//! byte/weight) and Q4 (½ byte/weight) weights read straight from their split +//! buffers, so the dense attention / shared-expert projections stay quantized- +//! resident instead of being pre-expanded to f16. The decode path is bandwidth- +//! bound on these projections (cold weights ≫ L2), so halving their bytes +//! roughly halves the per-layer cost. +//! +//! Variants: plain + `_coalesced` (contiguous-word warp walk, the decode fast +//! path) + fused `_relu2` (MoE up-proj activation) / `_accum` (router-weighted +//! down-proj into the layer accumulator), the `grouped_*` forms (one dispatch +//! for N row-groups each on their own x-slice), and the Q4 `_vec` / `_2row` +//! occupancy variants. The batched MoE expert-gather forms live in +//! `kernels/moe/gather_q4.rs`. Accumulation is f32 regardless of `T`. //! //! ## Q8_0 block (32 values) //! d (f16 scale) + 32 int8 quants; value[i] = d · q_i8[i] @@ -22,7 +28,7 @@ use metaltile::kernel; #[kernel] -pub fn ffai_gemv_q8( +pub fn mt_gemv_q8( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -63,7 +69,7 @@ pub fn ffai_gemv_q8( /// [1024,4096] Q8 slice, each on a different 4096-slice of the attention /// output) into a SINGLE dispatch instead of 8. #[kernel] -pub fn ffai_grouped_gemv_q8( +pub fn mt_grouped_gemv_q8( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -97,13 +103,13 @@ pub fn ffai_grouped_gemv_q8( } } -/// COALESCED per-token grouped Q8 gemv — same math as `ffai_grouped_gemv_q8` +/// COALESCED per-token grouped Q8 gemv — same math as `mt_grouped_gemv_q8` /// but the warp walks the row's `u32` words contiguously (lane j, j+32, …) so /// consecutive lanes hit consecutive addresses. The original strided by 8 u32 /// per lane (each lane owned a whole 32-int8 block), which only reached ~45% of /// DRAM bandwidth on GB10; this coalesced pattern is the decode-GEMV fast path. #[kernel] -pub fn ffai_gemv_q8_coalesced( +pub fn mt_gemv_q8_coalesced( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -142,7 +148,7 @@ pub fn ffai_gemv_q8_coalesced( /// Fuses a MoE expert's `up` projection and its activation into one dispatch /// (was gemv + a separate relu² kernel), keeping per-row occupancy. #[kernel] -pub fn ffai_gemv_q8_coalesced_relu2( +pub fn mt_gemv_q8_coalesced_relu2( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -185,7 +191,7 @@ pub fn ffai_gemv_q8_coalesced_relu2( /// device buffer (the router weight); loaded once per output row. #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_gemv_q8_coalesced_accum( +pub fn mt_gemv_q8_coalesced_accum( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -223,332 +229,6 @@ pub fn ffai_gemv_q8_coalesced_accum( } } -/// Copy a contiguous device slice `dst[i] = src[off + i]` — lets the Mamba -/// in_proj output be split (z / xBC / dt) ON-DEVICE instead of via a host -/// download, so the layer runs pure-async. `offbuf[0]` is the start offset. -#[kernel] -pub fn ffai_slice( - src: Tensor, - mut dst: Tensor, - #[constexpr] off: u32, - #[constexpr] len: u32, -) { - let i = program_id::<0>(); - if i < len { - store(dst[i], load(src[off + i])); - } -} - -/// Device dt for Mamba2: `dt[i] = softplus(dt_raw[i] + dt_bias[i])` (stable form). -/// Keeps the Mamba dt computation ON-DEVICE (no host round-trip). -#[kernel] -pub fn ffai_softplus_add( - a: Tensor, - b: Tensor, - mut out: Tensor, - #[constexpr] n: u32, -) { - let i = program_id::<0>(); - if i < n { - let x = load(a[i]) + load(b[i]); - let ax = select(x > 0.0f32, x, 0.0f32 - x); - let pos = select(x > 0.0f32, x, 0.0f32); - store(out[i], pos + log(1.0f32 + exp(0.0f32 - ax))); - } -} - -/// NemotronH/Zamba2 gated GROUPED RMSNorm (ON-DEVICE; removes the per-Mamba-layer -/// dl→host-norm→up sync). Gate-BEFORE-norm, per group of `gs`: g = y·silu(z); -/// out = g · rsqrt(mean_group(g²)+eps) · w. `y` fp32, z/w/out = T. One TG/group, -/// 4 elems/thread (block = gs/4), threadgroup reduce. -#[kernel] -pub fn ffai_gated_group_rmsnorm( - y: Tensor, - z: Tensor, - w: Tensor, - mut out: Tensor, - eps_buf: Tensor, - #[constexpr] gs: u32, -) { - let grp = program_id::<0>(); - let rs = grp * gs; - let col = tid * 4u32; - let in_bounds = col + 3u32 < gs; - let safe_col = select(in_bounds, col, 0u32); - let sb = rs + safe_col; - let y0 = load(y[sb]).cast::(); - let y1 = load(y[sb + 1u32]).cast::(); - let y2 = load(y[sb + 2u32]).cast::(); - let y3 = load(y[sb + 3u32]).cast::(); - let z0 = load(z[sb]).cast::(); - let z1 = load(z[sb + 1u32]).cast::(); - let z2 = load(z[sb + 2u32]).cast::(); - let z3 = load(z[sb + 3u32]).cast::(); - let g0 = y0 * (z0 / (1.0f32 + exp(0.0f32 - z0))); - let g1 = y1 * (z1 / (1.0f32 + exp(0.0f32 - z1))); - let g2 = y2 * (z2 / (1.0f32 + exp(0.0f32 - z2))); - let g3 = y3 * (z3 / (1.0f32 + exp(0.0f32 - z3))); - let raw = g0 * g0 + g1 * g1 + g2 * g2 + g3 * g3; - let partial = select(in_bounds, raw, 0.0f32); - let ssq = reduce_sum(partial); - let eps = load(eps_buf[0]); - let rms = rsqrt(ssq / (gs.cast::()) + eps); - if in_bounds { - let base = rs + col; - store(out[base], (g0 * rms * load(w[base]).cast::()).cast::()); - store(out[base + 1u32], (g1 * rms * load(w[base + 1u32]).cast::()).cast::()); - store(out[base + 2u32], (g2 * rms * load(w[base + 2u32]).cast::()).cast::()); - store(out[base + 3u32], (g3 * rms * load(w[base + 3u32]).cast::()).cast::()); - } -} - -/// 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 -/// router stays ON-DEVICE — no per-MoE-layer dl(gate)+host-topk+up(idx) sync round-trip. -#[kernel] -pub fn ffai_moe_sigmoid_bias( - logits: Tensor, - bias: Tensor, - mut unbiased: Tensor, - mut biased: Tensor, - #[constexpr] n: u32, -) { - let i = program_id::<0>(); - if i < n { - let s = 1.0f32 / (1.0f32 + exp(0.0f32 - load(logits[i]))); - store(unbiased[i], s); - store(biased[i], s + load(bias[i])); - } -} - -/// Scale a vector in place by a scalar (router weights × routed_scaling_factor). -#[kernel] -pub fn ffai_vscale(mut buf: Tensor, #[constexpr] scale: f32, #[constexpr] n: u32) { - let i = program_id::<0>(); - if i < n { - store(buf[i], load(buf[i]) * scale); - } -} - -/// Elementwise dtype cast f32 → f16. Compacts the attention KV cache to half -/// precision: at 32K context the sdpa read is bandwidth-bound, so halving the -/// cache bytes roughly halves the per-layer attention cost. One thread / elem. -#[kernel] -pub fn ffai_cast_f32_f16(src: Tensor, mut dst: Tensor, #[constexpr] n: u32) { - let i = program_id::<0>(); - if i < n { - store(dst[i], load(src[i]).cast::()); - } -} - -/// Elementwise dtype cast f16 → f32 (reverse): the sdpa f16 output is widened -/// back to f32 for the downstream o_proj Q4 GEMV, which consumes f32 activations. -#[kernel] -pub fn ffai_cast_f16_f32(src: Tensor, mut dst: Tensor, #[constexpr] n: u32) { - let i = program_id::<0>(); - if i < n { - store(dst[i], load(src[i]).cast::()); - } -} - -/// Roll a causal-conv state ON-DEVICE: `new = [old[conv_dim..], xbc]` (drop the -/// oldest conv_dim, append the current input) — keeps the Mamba conv history on -/// the GPU. `keep = (kc-2)*conv_dim`; indices clamped so both select branches -/// are in-bounds. -#[kernel] -pub fn ffai_conv_roll( - old: Tensor, - xbc: Tensor, - mut newst: Tensor, - #[constexpr] conv_dim: u32, - #[constexpr] keep: u32, - #[constexpr] n: u32, -) { - let i = program_id::<0>(); - if i < n { - let oi = select(i < keep, i + conv_dim, 0u32); - let xi = select(i < keep, 0u32, i - keep); - let v = select(i < keep, load(old[oi]), load(xbc[xi])); - store(newst[i], v); - } -} - -/// Batched MoE expert UP-projection + ReLU²: gathers the `top_k` selected -/// experts (indices in `idx`) from one contiguous `[n_exp*inter, hid]` Q4 weight -/// and computes all of them in ONE big GEMV — small per-expert matrices run at -/// ~52% DRAM bandwidth, but a [top_k*inter, hid] batch runs at ~90%. `out` is -/// `[top_k*inter]`. grid = top_k*inter threadgroups. -#[kernel] -pub fn ffai_moe_gather_q4_relu2( - qs: Tensor, - d_f32: Tensor, - x: Tensor, - idx: Tensor, - mut out: Tensor, - #[constexpr] k_in: u32, - #[constexpr] inter: u32, - #[constexpr] rows_per_tg: u32, -) { - // 2D grid [inter/rows_per_tg, top_k]: slot = tgid_y; `rows_per_tg` warps per - // TG each own one inter-row (multi-warp hides global-load latency, same as - // the dense gemv). rows_per_tg=1 is bit-identical (warp=0, lane=tid). - let warp = tid / 32u32; - let lane = tid % 32u32; - let local = tgid_x * rows_per_tg + warp; - let slot = tgid_y; - if local < inter { - let e = load(idx[slot]); - let row = e * inter + local; - let bpr = k_in / 32u32; - let nwords = bpr * 4u32; - let qs_base = row * bpr * 4u32; - let d_base = row * bpr; - let mut dot = 0.0f32; - for j in range(lane, nwords, 32u32) { - let block = j / 4u32; - let sub = j % 4u32; - let packed = load(qs[qs_base + j]); - let dd = load(d_f32[d_base + block]).cast::(); - let xb = block * 32u32 + sub * 8u32; - let mut blk = 0.0f32; - for i in range(0u32, 8u32, 1u32) { - let nib = (packed >> (i * 4u32)) & 0xfu32; - blk = blk - + (nib.cast::() - select(nib > 7u32, 16.0f32, 0.0f32)) - * load(x[xb + i]).cast::(); - } - dot = dot + dd * blk; - } - let total = simd_sum(dot); - if lane == 0u32 { - let rr = select(total > 0.0f32, total, 0.0f32); - store(out[slot * inter + local], (rr * rr).cast::()); - } - } -} - -/// Batched MoE expert DOWN-projection + router-weighted accumulate: for each -/// output row `h`, sums the `top_k` experts' `down[e,h]·x_slot` weighted by -/// `wts[slot]`, into `acc[h]`. One dispatch for all experts. `x` is the -/// `[top_k*inter]` up-relu² output; `qs` is the contiguous `[n_exp*hid, inter]`. -#[kernel] -#[allow(clippy::too_many_arguments)] -pub fn ffai_moe_gather_q4_down_accum( - qs: Tensor, - d_f32: Tensor, - x: Tensor, - idx: Tensor, - wts: Tensor, - mut acc: Tensor, - #[constexpr] inter: u32, - #[constexpr] hid: u32, - #[constexpr] top_k: u32, -) { - let h = tgid_x; - let lane = tid; - let bpr = inter / 32u32; - let nwords = bpr * 4u32; - let mut total = 0.0f32; - for slot in range(0u32, top_k, 1u32) { - let e = load(idx[slot]); - let row = e * hid + h; - let qs_base = row * bpr * 4u32; - let d_base = row * bpr; - let xoff = slot * inter; - let w = load(wts[slot]); - let mut dot = 0.0f32; - for j in range(lane, nwords, 32u32) { - let block = j / 4u32; - let sub = j % 4u32; - let packed = load(qs[qs_base + j]); - let dd = load(d_f32[d_base + block]).cast::(); - let xb = xoff + block * 32u32 + sub * 8u32; - for i in range(0u32, 8u32, 1u32) { - let nib = (packed >> (i * 4u32)) & 0xfu32; - dot = dot - + dd * (nib.cast::() - select(nib > 7u32, 16.0f32, 0.0f32)) - * load(x[xb + i]).cast::(); - } - } - total = total + w * simd_sum(dot); - } - if lane == 0u32 { - store(acc[h], (load(acc[h]).cast::() + total).cast::()); - } -} - -/// Batched MoE DOWN gather (no accumulate): `out[slot*hid + h] = down[e_slot, h]· -/// x_slot`, one big `[top_k*hid]` GEMV (grid top_k*hid ⇒ high occupancy, vs the -/// fused-accum variant's grid[hid] which serialized top_k experts at ~50% bw). -#[kernel] -pub fn ffai_moe_gather_q4_down( - qs: Tensor, - d_f32: Tensor, - x: Tensor, - idx: Tensor, - mut out: Tensor, - #[constexpr] inter: u32, - #[constexpr] hid: u32, - #[constexpr] rows_per_tg: u32, -) { - // 2D grid [hid/rows_per_tg, top_k]: rows_per_tg warps/TG, one hid-row each - // (multi-warp latency hiding). rows_per_tg=1 is bit-identical. - let warp = tid / 32u32; - let lane = tid % 32u32; - let local = tgid_x * rows_per_tg + warp; - let slot = tgid_y; - if local < hid { - let e = load(idx[slot]); - let row = e * hid + local; - let bpr = inter / 32u32; - let nwords = bpr * 4u32; - let qs_base = row * bpr * 4u32; - let d_base = row * bpr; - let xoff = slot * inter; - let mut dot = 0.0f32; - for j in range(lane, nwords, 32u32) { - let block = j / 4u32; - let sub = j % 4u32; - let packed = load(qs[qs_base + j]); - let dd = load(d_f32[d_base + block]).cast::(); - let xb = xoff + block * 32u32 + sub * 8u32; - let mut blk = 0.0f32; - for i in range(0u32, 8u32, 1u32) { - let nib = (packed >> (i * 4u32)) & 0xfu32; - blk = blk - + (nib.cast::() - select(nib > 7u32, 16.0f32, 0.0f32)) - * load(x[xb + i]).cast::(); - } - dot = dot + dd * blk; - } - let total = simd_sum(dot); - if lane == 0u32 { - store(out[slot * hid + local], total.cast::()); - } - } -} - -/// Router-weighted sum of the per-expert down outputs into `acc`: -/// `acc[h] += Σ_slot wts[slot]·downs[slot*hid + h]`. Cheap (grid hid). -#[kernel] -pub fn ffai_moe_weighted_sum( - downs: Tensor, - wts: Tensor, - mut acc: Tensor, - #[constexpr] hid: u32, - #[constexpr] top_k: u32, -) { - let h = program_id::<0>(); - if h < hid { - let mut t = load(acc[h]).cast::(); - for s in range(0u32, top_k, 1u32) { - t = t + load(wts[s]) * load(downs[s * hid + h]).cast::(); - } - store(acc[h], t.cast::()); - } -} - // ── Q4 (4-bit) coalesced gemv family — half the weight DRAM of Q8, the decode // bandwidth lever (decode reads cold weights: 35GB resident ≫ L2). Block 32, // symmetric int4 in [-7,7], one f32 scale/block. qs packs 8 nibbles per u32 @@ -556,7 +236,7 @@ pub fn ffai_moe_weighted_sum( /// Plain Q4 coalesced matvec: `out[r] = Σ_k dequant4(W[r,k]) · x[...]`. #[kernel] -pub fn ffai_gemv_q4_coalesced( +pub fn mt_gemv_q4_coalesced( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -609,7 +289,7 @@ pub fn ffai_gemv_q4_coalesced( /// stalls (ncu: the latency-bound GEMV's actual bottleneck). Coalesced: adjacent /// lanes read adjacent 16-byte blocks. #[kernel] -pub fn ffai_gemv_q4_vec( +pub fn mt_gemv_q4_vec( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -683,7 +363,7 @@ pub fn ffai_gemv_q4_vec( /// fall between the even/odd row pair. Odd `m_out` is safe: the dangling /// `row_b` clamps its weight reads to `row_a` and skips its store. #[kernel] -pub fn ffai_gemv_q4_coalesced_2row( +pub fn mt_gemv_q4_coalesced_2row( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -750,7 +430,7 @@ pub fn ffai_gemv_q4_coalesced_2row( /// global Q4 loads in flight to hide that latency. `rows_per_tg=1` is /// bit-identical to the original (warp=0, lane=tid, row=tgid_x). #[kernel] -pub fn ffai_gemv_q4_coalesced_relu2( +pub fn mt_gemv_q4_coalesced_relu2( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -800,7 +480,7 @@ pub fn ffai_gemv_q4_coalesced_relu2( /// hiding rationale as the relu2 variant; `rows_per_tg=1` is bit-identical. #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_gemv_q4_coalesced_accum( +pub fn mt_gemv_q4_coalesced_accum( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -847,34 +527,14 @@ pub fn ffai_gemv_q4_coalesced_accum( } } -/// Append a decode step's K (or V) into an IN-PLACE device KV cache, so the -/// growing context never round-trips through the host. `src` is `[nkv*hd]` (the -/// new token's per-head vectors), `dst` is the cache `[nkv, cap, hd]`, `posbuf[0]` -/// is the current position. Writes `dst[h, pos, :] = src[h, :]`. Runtime `pos` -/// rides in a buffer (NOT constexpr) so the kernel is compiled once, not per step. -#[kernel] -pub fn ffai_kv_append( - src: Tensor, - mut dst: Tensor, - posbuf: Tensor, - #[constexpr] hd: u32, - #[constexpr] cap: u32, -) { - let idx = program_id::<0>(); - let pos = load(posbuf[0]); - let h = idx / hd; - let dd = idx % hd; - store(dst[h * cap * hd + pos * hd + dd], load(src[idx])); -} - -/// BATCHED grouped Q8_0 gemv — ffai_grouped_gemv_q8 over `n_tokens` rows in +/// BATCHED grouped Q8_0 gemv — mt_grouped_gemv_q8 over `n_tokens` rows in /// ONE dispatch (grid z/y = token). Prefill O-LoRA looped the per-token /// grouped gemv N times; this folds it. x is [n_tokens, n_groups*k_in], /// out is [n_tokens, m_out]; n_groups = m_out/rows_per_group. /// Grid (Reduction): [m_out, n_tokens, 1], tg=[32,1,1]. #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_grouped_gemv_q8_rows( +pub fn mt_grouped_gemv_q8_rows( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -911,7 +571,7 @@ pub fn ffai_grouped_gemv_q8_rows( } /// TOKEN-TILED grouped Q8 gemv — the amortized fix for the prefill O-LoRA-A -/// hotspot. `ffai_grouped_gemv_q8_rows` re-reads each weight row from DRAM +/// hotspot. `mt_grouped_gemv_q8_rows` re-reads each weight row from DRAM /// once PER TOKEN (no amortization); at N=512 that's the single biggest /// op in the attention block (~47 ms/layer). Here each threadgroup owns one /// output row and a TILE of `tokens_per_tile` tokens: the Q8 weight block @@ -920,7 +580,7 @@ pub fn ffai_grouped_gemv_q8_rows( /// grid (threadgroups) = [m_out, ceil(n_tokens/T), 1], threadgroup [32,1,1]. #[kernel] #[allow(clippy::too_many_arguments)] -pub fn ffai_grouped_gemv_q8_rows_tiled( +pub fn mt_grouped_gemv_q8_rows_tiled( qs: Tensor, d_f32: Tensor, x: Tensor, @@ -982,10 +642,10 @@ pub mod kernel_benches { use metaltile::{bench, test::*}; use super::{ - ffai_gemv_q8, - ffai_grouped_gemv_q8, - ffai_grouped_gemv_q8_rows, - ffai_grouped_gemv_q8_rows_tiled, + mt_gemv_q8, + mt_grouped_gemv_q8, + mt_grouped_gemv_q8_rows, + mt_grouped_gemv_q8_rows_tiled, }; #[bench(dtypes = [f32, f16, bf16])] @@ -993,7 +653,7 @@ pub mod kernel_benches { let m_out = 4096usize; let k_in = 8192usize; let bpr = k_in / 32; - BenchSetup::new(ffai_gemv_q8::kernel_ir_for(dt)) + BenchSetup::new(mt_gemv_q8::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("qs", m_out * bpr * 8, DType::U32)) .buffer(BenchBuffer::random("d_f32", m_out * bpr, DType::F32)) @@ -1010,7 +670,7 @@ pub mod kernel_benches { let m_out = 8192usize; let k_in = 4096usize; let bpr = k_in / 32; - BenchSetup::new(ffai_grouped_gemv_q8::kernel_ir_for(dt)) + BenchSetup::new(mt_grouped_gemv_q8::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("qs", m_out * bpr * 8, DType::U32)) .buffer(BenchBuffer::random("d_f32", m_out * bpr, DType::F32)) @@ -1030,7 +690,7 @@ pub mod kernel_benches { let n_tokens = 256usize; let n_groups = m_out / 1024; let bpr = k_in / 32; - BenchSetup::new(ffai_grouped_gemv_q8_rows::kernel_ir_for(dt)) + BenchSetup::new(mt_grouped_gemv_q8_rows::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("qs", m_out * bpr * 8, DType::U32)) .buffer(BenchBuffer::random("d_f32", m_out * bpr, DType::F32)) @@ -1051,7 +711,7 @@ pub mod kernel_benches { let tokens_per_tile = 8usize; let n_groups = m_out / 1024; let bpr = k_in / 32; - BenchSetup::new(ffai_grouped_gemv_q8_rows_tiled::kernel_ir_for(dt)) + BenchSetup::new(mt_grouped_gemv_q8_rows_tiled::kernel_ir_for(dt)) .mode(KernelMode::Reduction) .buffer(BenchBuffer::random("qs", m_out * bpr * 8, DType::U32)) .buffer(BenchBuffer::random("d_f32", m_out * bpr, DType::F32)) diff --git a/crates/metaltile-std/src/kernels/gemm/mod.rs b/crates/metaltile-std/src/kernels/gemm/mod.rs index 4beb2545..3c7b5929 100644 --- a/crates/metaltile-std/src/kernels/gemm/mod.rs +++ b/crates/metaltile-std/src/kernels/gemm/mod.rs @@ -57,6 +57,9 @@ pub mod gemm_q4_mpp; pub mod gemm_q8; pub mod gemm_q8_mpp; +// Quantized inline-dequant GEMV (Q8_0 / Q4, coalesced + grouped/tiled forms). +pub mod gemv_quantized; + // Quantized patch-embed (block-scaled im2col + matmul). pub mod patch_embed_block_scaled; pub mod patch_embed_mma_block_scaled; diff --git a/crates/metaltile-std/src/kernels/kv_cache/cache.rs b/crates/metaltile-std/src/kernels/kv_cache/cache.rs index 8c818688..d3295e69 100644 --- a/crates/metaltile-std/src/kernels/kv_cache/cache.rs +++ b/crates/metaltile-std/src/kernels/kv_cache/cache.rs @@ -395,6 +395,27 @@ pub fn mt_bulk_dequant_kv_fp8_e5m2( store(out[dst_idx], w_real.cast::()); } +// Append one token's K/V into the cache (on-device, no host round-trip). +/// Append a decode step's K (or V) into an IN-PLACE device KV cache, so the +/// growing context never round-trips through the host. `src` is `[nkv*hd]` (the +/// new token's per-head vectors), `dst` is the cache `[nkv, cap, hd]`, `posbuf[0]` +/// is the current position. Writes `dst[h, pos, :] = src[h, :]`. Runtime `pos` +/// rides in a buffer (NOT constexpr) so the kernel is compiled once, not per step. +#[kernel] +pub fn mt_kv_append( + src: Tensor, + mut dst: Tensor, + posbuf: Tensor, + #[constexpr] hd: u32, + #[constexpr] cap: u32, +) { + let idx = program_id::<0>(); + let pos = load(posbuf[0]); + let h = idx / hd; + let dd = idx % hd; + store(dst[h * cap * hd + pos * hd + dd], load(src[idx])); +} + pub mod kernel_tests { use metaltile::{test::*, test_kernel}; diff --git a/crates/metaltile-std/src/kernels/mod.rs b/crates/metaltile-std/src/kernels/mod.rs index 7caa96d6..5977e79e 100644 --- a/crates/metaltile-std/src/kernels/mod.rs +++ b/crates/metaltile-std/src/kernels/mod.rs @@ -11,6 +11,7 @@ pub mod audio; pub mod convolution; pub mod gemm; pub mod kv_cache; +pub mod moe; pub mod norm; pub mod ops; pub mod rope; diff --git a/crates/metaltile-std/src/kernels/moe/gather_q4.rs b/crates/metaltile-std/src/kernels/moe/gather_q4.rs new file mode 100644 index 00000000..2d1d7305 --- /dev/null +++ b/crates/metaltile-std/src/kernels/moe/gather_q4.rs @@ -0,0 +1,184 @@ +//! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric +//! SPDX-License-Identifier: Apache-2.0 +//! Batched Q4 mixture-of-experts gather projections — fuse the per-expert +//! up / down GEMVs into single big dispatches (gather the top-k selected +//! experts from one contiguous Q4 weight, run them as one [top_k*rows, k] +//! GEMV at ~90% DRAM bandwidth instead of many small ~52% ones), plus the +//! router-weighted accumulate/sum that folds the expert outputs back in. + +use metaltile::kernel; + +/// Batched MoE expert UP-projection + ReLU²: gathers the `top_k` selected +/// experts (indices in `idx`) from one contiguous `[n_exp*inter, hid]` Q4 weight +/// and computes all of them in ONE big GEMV — small per-expert matrices run at +/// ~52% DRAM bandwidth, but a [top_k*inter, hid] batch runs at ~90%. `out` is +/// `[top_k*inter]`. grid = top_k*inter threadgroups. +#[kernel] +pub fn mt_moe_gather_q4_relu2( + qs: Tensor, + d_f32: Tensor, + x: Tensor, + idx: Tensor, + mut out: Tensor, + #[constexpr] k_in: u32, + #[constexpr] inter: u32, + #[constexpr] rows_per_tg: u32, +) { + // 2D grid [inter/rows_per_tg, top_k]: slot = tgid_y; `rows_per_tg` warps per + // TG each own one inter-row (multi-warp hides global-load latency, same as + // the dense gemv). rows_per_tg=1 is bit-identical (warp=0, lane=tid). + let warp = tid / 32u32; + let lane = tid % 32u32; + let local = tgid_x * rows_per_tg + warp; + let slot = tgid_y; + if local < inter { + let e = load(idx[slot]); + let row = e * inter + local; + let bpr = k_in / 32u32; + let nwords = bpr * 4u32; + let qs_base = row * bpr * 4u32; + let d_base = row * bpr; + let mut dot = 0.0f32; + for j in range(lane, nwords, 32u32) { + let block = j / 4u32; + let sub = j % 4u32; + let packed = load(qs[qs_base + j]); + let dd = load(d_f32[d_base + block]).cast::(); + let xb = block * 32u32 + sub * 8u32; + let mut blk = 0.0f32; + for i in range(0u32, 8u32, 1u32) { + let nib = (packed >> (i * 4u32)) & 0xfu32; + blk = blk + + (nib.cast::() - select(nib > 7u32, 16.0f32, 0.0f32)) + * load(x[xb + i]).cast::(); + } + dot = dot + dd * blk; + } + let total = simd_sum(dot); + if lane == 0u32 { + let rr = select(total > 0.0f32, total, 0.0f32); + store(out[slot * inter + local], (rr * rr).cast::()); + } + } +} + +/// Batched MoE expert DOWN-projection + router-weighted accumulate: for each +/// output row `h`, sums the `top_k` experts' `down[e,h]·x_slot` weighted by +/// `wts[slot]`, into `acc[h]`. One dispatch for all experts. `x` is the +/// `[top_k*inter]` up-relu² output; `qs` is the contiguous `[n_exp*hid, inter]`. +#[kernel] +#[allow(clippy::too_many_arguments)] +pub fn mt_moe_gather_q4_down_accum( + qs: Tensor, + d_f32: Tensor, + x: Tensor, + idx: Tensor, + wts: Tensor, + mut acc: Tensor, + #[constexpr] inter: u32, + #[constexpr] hid: u32, + #[constexpr] top_k: u32, +) { + let h = tgid_x; + let lane = tid; + let bpr = inter / 32u32; + let nwords = bpr * 4u32; + let mut total = 0.0f32; + for slot in range(0u32, top_k, 1u32) { + let e = load(idx[slot]); + let row = e * hid + h; + let qs_base = row * bpr * 4u32; + let d_base = row * bpr; + let xoff = slot * inter; + let w = load(wts[slot]); + let mut dot = 0.0f32; + for j in range(lane, nwords, 32u32) { + let block = j / 4u32; + let sub = j % 4u32; + let packed = load(qs[qs_base + j]); + let dd = load(d_f32[d_base + block]).cast::(); + let xb = xoff + block * 32u32 + sub * 8u32; + for i in range(0u32, 8u32, 1u32) { + let nib = (packed >> (i * 4u32)) & 0xfu32; + dot = dot + + dd * (nib.cast::() - select(nib > 7u32, 16.0f32, 0.0f32)) + * load(x[xb + i]).cast::(); + } + } + total = total + w * simd_sum(dot); + } + if lane == 0u32 { + store(acc[h], (load(acc[h]).cast::() + total).cast::()); + } +} + +/// Batched MoE DOWN gather (no accumulate): `out[slot*hid + h] = down[e_slot, h]· +/// x_slot`, one big `[top_k*hid]` GEMV (grid top_k*hid ⇒ high occupancy, vs the +/// fused-accum variant's grid[hid] which serialized top_k experts at ~50% bw). +#[kernel] +pub fn mt_moe_gather_q4_down( + qs: Tensor, + d_f32: Tensor, + x: Tensor, + idx: Tensor, + mut out: Tensor, + #[constexpr] inter: u32, + #[constexpr] hid: u32, + #[constexpr] rows_per_tg: u32, +) { + // 2D grid [hid/rows_per_tg, top_k]: rows_per_tg warps/TG, one hid-row each + // (multi-warp latency hiding). rows_per_tg=1 is bit-identical. + let warp = tid / 32u32; + let lane = tid % 32u32; + let local = tgid_x * rows_per_tg + warp; + let slot = tgid_y; + if local < hid { + let e = load(idx[slot]); + let row = e * hid + local; + let bpr = inter / 32u32; + let nwords = bpr * 4u32; + let qs_base = row * bpr * 4u32; + let d_base = row * bpr; + let xoff = slot * inter; + let mut dot = 0.0f32; + for j in range(lane, nwords, 32u32) { + let block = j / 4u32; + let sub = j % 4u32; + let packed = load(qs[qs_base + j]); + let dd = load(d_f32[d_base + block]).cast::(); + let xb = xoff + block * 32u32 + sub * 8u32; + let mut blk = 0.0f32; + for i in range(0u32, 8u32, 1u32) { + let nib = (packed >> (i * 4u32)) & 0xfu32; + blk = blk + + (nib.cast::() - select(nib > 7u32, 16.0f32, 0.0f32)) + * load(x[xb + i]).cast::(); + } + dot = dot + dd * blk; + } + let total = simd_sum(dot); + if lane == 0u32 { + store(out[slot * hid + local], total.cast::()); + } + } +} + +/// Router-weighted sum of the per-expert down outputs into `acc`: +/// `acc[h] += Σ_slot wts[slot]·downs[slot*hid + h]`. Cheap (grid hid). +#[kernel] +pub fn mt_moe_weighted_sum( + downs: Tensor, + wts: Tensor, + mut acc: Tensor, + #[constexpr] hid: u32, + #[constexpr] top_k: u32, +) { + let h = program_id::<0>(); + if h < hid { + let mut t = load(acc[h]).cast::(); + for s in range(0u32, top_k, 1u32) { + t = t + load(wts[s]) * load(downs[s * hid + h]).cast::(); + } + store(acc[h], t.cast::()); + } +} diff --git a/crates/metaltile-std/src/kernels/moe/mod.rs b/crates/metaltile-std/src/kernels/moe/mod.rs new file mode 100644 index 00000000..9e90d8e7 --- /dev/null +++ b/crates/metaltile-std/src/kernels/moe/mod.rs @@ -0,0 +1,10 @@ +//! 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. + +pub mod gather_q4; +pub mod sigmoid_bias; diff --git a/crates/metaltile-std/src/kernels/moe/sigmoid_bias.rs b/crates/metaltile-std/src/kernels/moe/sigmoid_bias.rs new file mode 100644 index 00000000..9a2ccb33 --- /dev/null +++ b/crates/metaltile-std/src/kernels/moe/sigmoid_bias.rs @@ -0,0 +1,26 @@ +//! Copyright 2026 0xClandestine, Ekryski, TheTom, Ambisphaeric +//! SPDX-License-Identifier: Apache-2.0 +//! MoE router pre-score (sigmoid + correction-bias), kept on-device so the +//! whole router (pre-score -> top-k -> gather) avoids a host round-trip. + +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 +/// router stays ON-DEVICE — no per-MoE-layer dl(gate)+host-topk+up(idx) sync round-trip. +#[kernel] +pub fn mt_moe_sigmoid_bias( + logits: Tensor, + bias: Tensor, + mut unbiased: Tensor, + mut biased: Tensor, + #[constexpr] n: u32, +) { + let i = program_id::<0>(); + if i < n { + let s = 1.0f32 / (1.0f32 + exp(0.0f32 - load(logits[i]))); + store(unbiased[i], s); + store(biased[i], s + load(bias[i])); + } +} diff --git a/crates/metaltile-std/src/kernels/ops/copy.rs b/crates/metaltile-std/src/kernels/ops/copy.rs index 2291c4ba..feef602b 100644 --- a/crates/metaltile-std/src/kernels/ops/copy.rs +++ b/crates/metaltile-std/src/kernels/ops/copy.rs @@ -11,6 +11,23 @@ pub fn mt_copy(a: Tensor, out: Tensor) { } /// New-syntax correctness for `mt_copy` (elementwise, bit-exact). +// Contiguous device-slice copy (split a packed buffer on-device). +/// Copy a contiguous device slice `dst[i] = src[off + i]` — lets the Mamba +/// in_proj output be split (z / xBC / dt) ON-DEVICE instead of via a host +/// download, so the layer runs pure-async. `offbuf[0]` is the start offset. +#[kernel] +pub fn mt_slice( + src: Tensor, + mut dst: Tensor, + #[constexpr] off: u32, + #[constexpr] len: u32, +) { + let i = program_id::<0>(); + if i < len { + store(dst[i], load(src[off + i])); + } +} + pub mod kernel_tests { use metaltile::{test::*, test_kernel}; diff --git a/crates/metaltile-std/src/kernels/ops/unary.rs b/crates/metaltile-std/src/kernels/ops/unary.rs index c0632cbf..88c5ccdc 100644 --- a/crates/metaltile-std/src/kernels/ops/unary.rs +++ b/crates/metaltile-std/src/kernels/ops/unary.rs @@ -237,6 +237,37 @@ pub fn mt_erfinv(a: Tensor, out: Tensor) { /// generous-but-bounded per-dtype band that still catches an empty-body or /// wrong-formula kernel. `erf`/`gelu`/`erfinv` are bench-only — there's no /// std f32 oracle for them (the legacy test didn't cover them either). +// Device-glue casts + scalar scale (keep activation pipelines on-GPU). +/// Scale a vector in place by a scalar (router weights × routed_scaling_factor). +#[kernel] +pub fn mt_vscale(mut buf: Tensor, #[constexpr] scale: f32, #[constexpr] n: u32) { + let i = program_id::<0>(); + if i < n { + store(buf[i], load(buf[i]) * scale); + } +} + +/// Elementwise dtype cast f32 → f16. Compacts the attention KV cache to half +/// precision: at 32K context the sdpa read is bandwidth-bound, so halving the +/// cache bytes roughly halves the per-layer attention cost. One thread / elem. +#[kernel] +pub fn mt_cast_f32_f16(src: Tensor, mut dst: Tensor, #[constexpr] n: u32) { + let i = program_id::<0>(); + if i < n { + store(dst[i], load(src[i]).cast::()); + } +} + +/// Elementwise dtype cast f16 → f32 (reverse): the sdpa f16 output is widened +/// back to f32 for the downstream o_proj Q4 GEMV, which consumes f32 activations. +#[kernel] +pub fn mt_cast_f16_f32(src: Tensor, mut dst: Tensor, #[constexpr] n: u32) { + let i = program_id::<0>(); + if i < n { + store(dst[i], load(src[i]).cast::()); + } +} + pub mod kernel_tests { use metaltile::{core::ir::Kernel, test::*, test_kernel}; diff --git a/crates/metaltile-std/src/kernels/ssm/scan.rs b/crates/metaltile-std/src/kernels/ssm/scan.rs index 068c0bdd..33948288 100644 --- a/crates/metaltile-std/src/kernels/ssm/scan.rs +++ b/crates/metaltile-std/src/kernels/ssm/scan.rs @@ -229,6 +229,65 @@ pub fn mt_ssm_step_grouped( } } +// Mamba on-device dt + gated grouped RMSNorm (elementwise / one-TG forms). +/// Device dt for Mamba2: `dt[i] = softplus(dt_raw[i] + dt_bias[i])` (stable form). +/// Keeps the Mamba dt computation ON-DEVICE (no host round-trip). +#[kernel] +pub fn mt_softplus_add(a: Tensor, b: Tensor, mut out: Tensor, #[constexpr] n: u32) { + let i = program_id::<0>(); + if i < n { + let x = load(a[i]) + load(b[i]); + let ax = select(x > 0.0f32, x, 0.0f32 - x); + let pos = select(x > 0.0f32, x, 0.0f32); + store(out[i], pos + log(1.0f32 + exp(0.0f32 - ax))); + } +} + +/// NemotronH/Zamba2 gated GROUPED RMSNorm (ON-DEVICE; removes the per-Mamba-layer +/// dl→host-norm→up sync). Gate-BEFORE-norm, per group of `gs`: g = y·silu(z); +/// out = g · rsqrt(mean_group(g²)+eps) · w. `y` fp32, z/w/out = T. One TG/group, +/// 4 elems/thread (block = gs/4), threadgroup reduce. +#[kernel] +pub fn mt_gated_group_rmsnorm( + y: Tensor, + z: Tensor, + w: Tensor, + mut out: Tensor, + eps_buf: Tensor, + #[constexpr] gs: u32, +) { + let grp = program_id::<0>(); + let rs = grp * gs; + let col = tid * 4u32; + let in_bounds = col + 3u32 < gs; + let safe_col = select(in_bounds, col, 0u32); + let sb = rs + safe_col; + let y0 = load(y[sb]).cast::(); + let y1 = load(y[sb + 1u32]).cast::(); + let y2 = load(y[sb + 2u32]).cast::(); + let y3 = load(y[sb + 3u32]).cast::(); + let z0 = load(z[sb]).cast::(); + let z1 = load(z[sb + 1u32]).cast::(); + let z2 = load(z[sb + 2u32]).cast::(); + let z3 = load(z[sb + 3u32]).cast::(); + let g0 = y0 * (z0 / (1.0f32 + exp(0.0f32 - z0))); + let g1 = y1 * (z1 / (1.0f32 + exp(0.0f32 - z1))); + let g2 = y2 * (z2 / (1.0f32 + exp(0.0f32 - z2))); + let g3 = y3 * (z3 / (1.0f32 + exp(0.0f32 - z3))); + let raw = g0 * g0 + g1 * g1 + g2 * g2 + g3 * g3; + let partial = select(in_bounds, raw, 0.0f32); + let ssq = reduce_sum(partial); + let eps = load(eps_buf[0]); + let rms = rsqrt(ssq / (gs.cast::()) + eps); + if in_bounds { + let base = rs + col; + store(out[base], (g0 * rms * load(w[base]).cast::()).cast::()); + store(out[base + 1u32], (g1 * rms * load(w[base + 1u32]).cast::()).cast::()); + store(out[base + 2u32], (g2 * rms * load(w[base + 2u32]).cast::()).cast::()); + store(out[base + 3u32], (g3 * rms * load(w[base + 3u32]).cast::()).cast::()); + } +} + pub mod kernel_tests { use metaltile::{test::*, test_kernel}; diff --git a/crates/metaltile-std/tests/gemv_q8_correctness.rs b/crates/metaltile-std/tests/gemv_q8_correctness.rs index bde5c5e6..94597ca2 100644 --- a/crates/metaltile-std/tests/gemv_q8_correctness.rs +++ b/crates/metaltile-std/tests/gemv_q8_correctness.rs @@ -1,6 +1,6 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_gemv_q8` — Q8_0 inline-dequant gemv +//! GPU correctness for `mt_gemv_q8` — Q8_0 inline-dequant gemv //! vs a CPU reference using the same dequant (value = d * int8). #![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::gemv_q8::ffai_gemv_q8; +use metaltile_std::kernels::gemm::gemv_quantized::mt_gemv_q8; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; @@ -59,7 +59,7 @@ fn run_case(dt: Dt, k_in: usize, m_out: usize, tol: f32) { buffers.insert("m_out".into(), (m_out as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("ctx"); - let mut kernel = ffai_gemv_q8::kernel_ir_for(dt.to_dtype()); + let mut kernel = mt_gemv_q8::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]) @@ -86,7 +86,7 @@ fn gemv_q8_f16() { #[test] fn grouped_gemv_q8_f32() { - use metaltile_std::ffai::gemv_q8::ffai_grouped_gemv_q8; + use metaltile_std::kernels::gemm::gemv_quantized::mt_grouped_gemv_q8; let _g = gpu_lock(); let k_in = 4096usize; let rows_per_group = 1024usize; @@ -125,7 +125,7 @@ fn grouped_gemv_q8_f32() { buffers.insert("m_out".into(), (m_out as u32).to_le_bytes().to_vec()); buffers.insert("rows_per_group".into(), (rows_per_group as u32).to_le_bytes().to_vec()); let ctx = Context::new().expect("ctx"); - let mut kernel = ffai_grouped_gemv_q8::kernel_ir_for(Dt::F32.to_dtype()); + let mut kernel = mt_grouped_gemv_q8::kernel_ir_for(Dt::F32.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/grouped_gemv_q8_rows_correctness.rs b/crates/metaltile-std/tests/grouped_gemv_q8_rows_correctness.rs index ad9ec0dc..51f4556f 100644 --- a/crates/metaltile-std/tests/grouped_gemv_q8_rows_correctness.rs +++ b/crates/metaltile-std/tests/grouped_gemv_q8_rows_correctness.rs @@ -1,7 +1,7 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! Batched grouped Q8 gemv (ffai_grouped_gemv_q8_rows) must equal the -//! per-token single kernel (ffai_grouped_gemv_q8) row by row. NO model load. +//! Batched grouped Q8 gemv (mt_grouped_gemv_q8_rows) must equal the +//! per-token single kernel (mt_grouped_gemv_q8) row by row. NO model load. #![cfg(target_os = "macos")] mod common; @@ -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::gemv_q8::{ffai_grouped_gemv_q8, ffai_grouped_gemv_q8_rows}; +use metaltile_std::kernels::gemm::gemv_quantized::{mt_grouped_gemv_q8, mt_grouped_gemv_q8_rows}; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; @@ -58,7 +58,7 @@ fn grouped_gemv_q8_rows_matches_single() { // Reference: single kernel per token. let mut want = vec![0.0f32; n_tokens * m_out]; - let mut ks = ffai_grouped_gemv_q8::kernel_ir_for(Dt::F32.to_dtype()); + let mut ks = mt_grouped_gemv_q8::kernel_ir_for(Dt::F32.to_dtype()); ks.mode = KernelMode::Reduction; for t in 0..n_tokens { let xt = &x[t * n_groups * k_in..(t + 1) * n_groups * k_in]; @@ -79,7 +79,7 @@ fn grouped_gemv_q8_rows_matches_single() { bb.insert("d_f32".into(), pack_bytes(&d, Dt::F32)); bb.insert("x".into(), pack_bytes(&x, Dt::F32)); bb.insert("out".into(), pack_bytes(&vec![0.0f32; n_tokens * m_out], Dt::F32)); - let mut kr = ffai_grouped_gemv_q8_rows::kernel_ir_for(Dt::F32.to_dtype()); + let mut kr = mt_grouped_gemv_q8_rows::kernel_ir_for(Dt::F32.to_dtype()); kr.mode = KernelMode::Reduction; let r = ctx .dispatch_with_grid(&kr, &bb, &BTreeMap::new(), [m_out, n_tokens, 1], [32, 1, 1]) diff --git a/crates/metaltile-std/tests/grouped_gemv_q8_rows_tiled_correctness.rs b/crates/metaltile-std/tests/grouped_gemv_q8_rows_tiled_correctness.rs index c0f6e49e..5cebfe99 100644 --- a/crates/metaltile-std/tests/grouped_gemv_q8_rows_tiled_correctness.rs +++ b/crates/metaltile-std/tests/grouped_gemv_q8_rows_tiled_correctness.rs @@ -1,8 +1,8 @@ //! Copyright 2026 TheTom //! SPDX-License-Identifier: Apache-2.0 -//! GPU correctness for `ffai::ffai_grouped_gemv_q8_rows_tiled` — the +//! GPU correctness for `mt_grouped_gemv_q8_rows_tiled` — the //! token-TILED grouped Q8 gemv (8-fold weight-DRAM amortization). It must -//! equal the proven `ffai_grouped_gemv_q8_rows` row-by-row: same grouped Q8 +//! equal the proven `mt_grouped_gemv_q8_rows` row-by-row: same grouped Q8 //! dequant and dot product, only the per-token weight reuse differs. NO model //! load. (Mirrors `grouped_gemv_q8_rows_correctness.rs`, which validates //! `_rows` against the per-token single kernel.) @@ -14,7 +14,10 @@ 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::gemv_q8::{ffai_grouped_gemv_q8_rows, ffai_grouped_gemv_q8_rows_tiled}; +use metaltile_std::kernels::gemm::gemv_quantized::{ + mt_grouped_gemv_q8_rows, + mt_grouped_gemv_q8_rows_tiled, +}; fn xorshift(s: &mut u32) -> u32 { let mut x = *s; @@ -64,7 +67,7 @@ fn grouped_gemv_q8_rows_tiled_matches_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("rows_per_group".into(), (rows_per_group as u32).to_le_bytes().to_vec()); - let mut kr = ffai_grouped_gemv_q8_rows::kernel_ir_for(Dt::F32.to_dtype()); + let mut kr = mt_grouped_gemv_q8_rows::kernel_ir_for(Dt::F32.to_dtype()); kr.mode = KernelMode::Reduction; let rr = ctx .dispatch_with_grid(&kr, &bref, &BTreeMap::new(), [m_out, n_tokens, 1], [32, 1, 1]) @@ -81,7 +84,7 @@ fn grouped_gemv_q8_rows_tiled_matches_rows() { bt.insert("m_out".into(), (m_out as u32).to_le_bytes().to_vec()); bt.insert("rows_per_group".into(), (rows_per_group as u32).to_le_bytes().to_vec()); bt.insert("n_tokens".into(), (n_tokens as u32).to_le_bytes().to_vec()); - let mut kt = ffai_grouped_gemv_q8_rows_tiled::kernel_ir_for(Dt::F32.to_dtype()); + let mut kt = mt_grouped_gemv_q8_rows_tiled::kernel_ir_for(Dt::F32.to_dtype()); kt.mode = KernelMode::Reduction; let gy = n_tokens.div_ceil(8); let rt = diff --git a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md index 53b8d91f..62dbdb79 100644 --- a/docs/specs/KERNEL_CONSOLIDATION_PLAN.md +++ b/docs/specs/KERNEL_CONSOLIDATION_PLAN.md @@ -45,26 +45,31 @@ the FFAI emit path is unaffected). crates/metaltile-std/src/kernels/ ops/ ✅ DONE — elementwise/core primitives: binary · unary · ternary · copy · arange · random · reduce · arg_reduce · scan · indexing · gather/scatter · hadamard · - fence · clamp · logsumexp · vector_add · axpy · strided · gated_activation + fence · clamp · logsumexp · vector_add · axpy · strided · gated_activation · + slice · vscale · cast(f32↔f16) gemm/ ✅ DONE — dense: gemm · gemv(_masked,_axpy) · patch_embed(_mma) · steel/gemm; quantized: quantized_*(+mpp/nax/int8/dynamic_m) · fp_quantized_* · block_scaled_* · dequant_gemv · gemm_q8(_mpp)/q4_mpp · batched_{qkv,4}(_block_scaled)_{qgemv,qmm} - · patch_embed(_mma)_block_scaled (same folder; format-axis fold deferred, §7) + · patch_embed(_mma)_block_scaled · gemv_quantized (Q8/Q4 inline-dequant gemv, + ex-gemv_q8 grab-bag) (same folder; format-axis fold deferred, §7) sdpa/ ALL attention: bidirectional(+relpos/windowed/conformer) · decode(+d64..d512/ 2pass/batched/sink) · multi(+d256/tree-mask) · prefill_mma · flash_quantized · aura_flash · steel/attn - moe/ moe orchestration · mpp(bm8/bm64 × int8) · bgemm/gemv(q2k/iq2xxs) · block_scaled_moe + 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 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 (see §4) + convolution/ ✅ DONE — conv1d/2d/3d · depthwise · winograd · steel_conv · conv1d_causal(_roll) (see §4) ssm/ ✅ DONE — ssm(_replay) · gated_delta(+wy/prep/chunk) · mamba pregate-rmsnorm + (gated_group_rmsnorm(_batched)) · softplus_add(_rows) quant/ INFRA + the op×format matrix (§7): codec · format · gguf · block_scaled_* · quantized_* · fp_quantized_* · affine · aura codec stack · dequant_* audio/ ✅ DONE — mel_spectrogram(+magnitude/stft/filterbank) · lstm · vocoder · snake1d · upsample vision/ ✅ DONE — resize_normalize(+bicubic) · im2col · patch_unfold · pos_emb_2d · avg_pool2d · transpose_th · frame_diff · broadcast_affine sampling/ ✅ DONE — logits_topk/top_p/min_p/processors · categorical_sample · softmax · sort - kv_cache/ ✅ DONE — kv_cache(_update_many) · fft + kv_cache/ ✅ DONE — kv_cache(_update_many) · kv_append · fft primitives.rs cross-family decode/reduce ops (mt_decode_e2m1/e4m3/e5m2/e8m0, mt_unpack_nbit, …) mod.rs pub mod ops; pub mod gemm; pub mod sdpa; … ```