From b9d8629280a679cae703bc25da004e577c2bd76c Mon Sep 17 00:00:00 2001 From: Tom Turney Date: Thu, 25 Jun 2026 13:58:58 -0500 Subject: [PATCH 1/2] test: gemv_q4 oracle uses f16 scales (match the kernel) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit mt_gemv_q4_coalesced reads f16 scales (resident-weight decode/prefill feed f16). gemv_q4_matches_cpu_dequant was uploading the quantize_q4 f32 scales and referencing against f32 — feeding the wrong dtype and motivating the ffai-kernels#33 regression. Upload the scales as f16 and reference against the f16-rounded value so the kernel and CPU oracle agree exactly. Pairs with ffai-kernels#36 (reverts the kernel to f16). Merge that first. --- .../backends/ffai-metal/tests/gemv_q8.rs | 32 ++++++++++++++++++- 1 file changed, 31 insertions(+), 1 deletion(-) diff --git a/rust/crates/backends/ffai-metal/tests/gemv_q8.rs b/rust/crates/backends/ffai-metal/tests/gemv_q8.rs index 651e8813..8497ed78 100644 --- a/rust/crates/backends/ffai-metal/tests/gemv_q8.rs +++ b/rust/crates/backends/ffai-metal/tests/gemv_q8.rs @@ -10,6 +10,34 @@ use ffai_metal::MetalDevice; fn tb_f32(v: &[f32]) -> Vec { v.iter().flat_map(|x| x.to_le_bytes()).collect() } fn tb_u32(v: &[u32]) -> Vec { v.iter().flat_map(|x| x.to_le_bytes()).collect() } fn fb(b: &[u8]) -> Vec { b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect() } +// f32<->f16 for the q4 coalesced kernel, whose scales are f16 (see ffai-kernels +// gemv_quantized: mt_gemv_q4_coalesced reads f16 scales). Feed f16 + reference +// against the f16-rounded scale so kernel and CPU oracle agree exactly. +fn f32_to_f16(f: f32) -> u16 { + let x = f.to_bits(); + let sign = ((x >> 16) & 0x8000) as u16; + let e = ((x >> 23) & 0xff) as i32 - 112; + if e <= 0 { return sign; } + if e >= 0x1f { return sign | 0x7c00; } + let m = (x >> 13) & 0x3ff; + let round = (x >> 12) & 1; + sign | ((((e as u32) << 10) | m) + round) as u16 +} +fn f16_to_f32(h: u16) -> f32 { + let sign = (h as u32 & 0x8000) << 16; + let exp = (h as u32 >> 10) & 0x1f; + let mant = h as u32 & 0x3ff; + let bits = if exp == 0 { + if mant == 0 { sign } else { + let mut e = -1i32; let mut m = mant; + while m & 0x400 == 0 { m <<= 1; e -= 1; } m &= 0x3ff; + sign | (((e + 127 - 14) as u32) << 23) | (m << 13) + } + } else if exp == 0x1f { sign | 0x7f80_0000 | (mant << 13) } + else { sign | ((exp + 112) << 23) | (mant << 13) }; + f32::from_bits(bits) +} +fn tb_f16(v: &[f32]) -> Vec { v.iter().flat_map(|&f| f32_to_f16(f).to_le_bytes()).collect() } #[test] fn gemv_q8_matches_cpu_dequant_dot() { @@ -98,6 +126,8 @@ fn gemv_q4_matches_cpu_dequant() { let w: Vec = (0..m * k).map(|_| rng()).collect(); let x: Vec = (0..k).map(|_| rng()).collect(); let (qs, sc) = ffai_ops::quantize_q4(&w, m, k); + // kernel uses f16 scales -> reference against the f16-rounded value + let sc: Vec = sc.iter().map(|&v| f16_to_f32(f32_to_f16(v))).collect(); let bpr = k / 32; let mut want = vec![0f32; m]; for r in 0..m { let mut a = 0f32; @@ -106,7 +136,7 @@ fn gemv_q4_matches_cpu_dequant() { let q = nib as i32 - if nib > 7 { 16 } else { 0 }; a += dd * q as f32 * x[b*32+i]; } } want[r] = a; } let qt = Tensor::new(d.upload(&tb_u32(&qs)).unwrap(), vec![qs.len()], DType::U32); - let st = Tensor::new(d.upload(&tb_f32(&sc)).unwrap(), vec![sc.len()], DType::F32); + let st = Tensor::new(d.upload(&tb_f16(&sc)).unwrap(), vec![sc.len()], DType::F16); let xt = Tensor::new(d.upload(&tb_f32(&x)).unwrap(), vec![k], DType::F32); let ot = ffai_ops::gemv_q4(d, &qt, &st, &xt, m, k, m).unwrap(); let mut ob = vec![0u8; m*4]; d.download(ot.buffer.as_ref(), &mut ob).unwrap(); From 576530564de1db31e54e4a36598571aefb768857 Mon Sep 17 00:00:00 2001 From: Tom Turney Date: Sat, 18 Jul 2026 08:22:53 -0500 Subject: [PATCH 2/2] fix(rust): resync rust/ tree against renamed ffai-kernels dev MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Cargo.lock pinned pre-rename kernels (d4cafc60), so the tree still compiled against MetalTileError and mt_-prefixed kernel fns; against current dev the Rust CI job fails at the first unresolved import. - Rename sweep: mt_ -> ffai_ identifiers and kernel-name references across rust/crates (7 files), matching kernels #40/#41. - Repin the six ffai-kernels crates to dev 9fe86104. Audited all ffai-ops dispatch sites against current kernel definitions (param order, constexprs, grid/TG geometry, dispatch-invariant docs) — no geometry drift; the intermittent local gemv failures are the known parallel PSO-compile race the Rust CI job already pins to one test thread. cargo check --workspace --all-targets clean; cargo test -p ffai-metal -p ffai-ops (single-thread, on GPU): 47 passed, 0 failed, 1 ignored. --- rust/Cargo.lock | 12 +- rust/crates/backends/ffai-cuda/src/imp.rs | 4 +- .../backends/ffai-cuda/tests/cuda_smoke.rs | 14 +- .../backends/ffai-cuda/tests/ssm_test.rs | 6 +- rust/crates/backends/ffai-metal/src/lib.rs | 4 +- .../backends/ffai-metal/tests/gemv_q8.rs | 2 +- .../backends/ffai-metal/tests/ssm_test.rs | 6 +- rust/crates/backends/ffai-vulkan/src/imp.rs | 4 +- rust/crates/ffai-modeltests/src/lib.rs | 12 +- rust/crates/ffai-ops/src/lib.rs | 188 +++++++++--------- rust/crates/ffai-ops/src/ssd_scan_portable.rs | 20 +- 11 files changed, 136 insertions(+), 136 deletions(-) diff --git a/rust/Cargo.lock b/rust/Cargo.lock index adb4b071..b428248b 100644 --- a/rust/Cargo.lock +++ b/rust/Cargo.lock @@ -138,7 +138,7 @@ dependencies = [ [[package]] name = "ffai-kernels" version = "0.1.0" -source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#d4cafc60ea4a43b411fdec98ede42cfeec3cd49d" +source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#9fe8610439fda6ffee46c9dd60fae3b0c7ecd1ec" dependencies = [ "bytemuck", "ffai-kernels-codegen", @@ -157,7 +157,7 @@ dependencies = [ [[package]] name = "ffai-kernels-codegen" version = "0.1.0" -source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#d4cafc60ea4a43b411fdec98ede42cfeec3cd49d" +source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#9fe8610439fda6ffee46c9dd60fae3b0c7ecd1ec" dependencies = [ "ffai-kernels-core", "inventory", @@ -172,7 +172,7 @@ dependencies = [ [[package]] name = "ffai-kernels-core" version = "0.1.0" -source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#d4cafc60ea4a43b411fdec98ede42cfeec3cd49d" +source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#9fe8610439fda6ffee46c9dd60fae3b0c7ecd1ec" dependencies = [ "ffai-kernels-macros", "inventory", @@ -186,7 +186,7 @@ dependencies = [ [[package]] name = "ffai-kernels-macros" version = "0.1.0" -source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#d4cafc60ea4a43b411fdec98ede42cfeec3cd49d" +source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#9fe8610439fda6ffee46c9dd60fae3b0c7ecd1ec" dependencies = [ "proc-macro2", "quote", @@ -196,7 +196,7 @@ dependencies = [ [[package]] name = "ffai-kernels-runtime" version = "0.1.0" -source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#d4cafc60ea4a43b411fdec98ede42cfeec3cd49d" +source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#9fe8610439fda6ffee46c9dd60fae3b0c7ecd1ec" dependencies = [ "ffai-kernels-codegen", "ffai-kernels-core", @@ -213,7 +213,7 @@ dependencies = [ [[package]] name = "ffai-kernels-std" version = "0.1.0" -source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#d4cafc60ea4a43b411fdec98ede42cfeec3cd49d" +source = "git+https://github.com/thewafflehaus/ffai-kernels?branch=dev#9fe8610439fda6ffee46c9dd60fae3b0c7ecd1ec" dependencies = [ "ffai-kernels", "half", diff --git a/rust/crates/backends/ffai-cuda/src/imp.rs b/rust/crates/backends/ffai-cuda/src/imp.rs index 6e9d49c5..db90f6c8 100644 --- a/rust/crates/backends/ffai-cuda/src/imp.rs +++ b/rust/crates/backends/ffai-cuda/src/imp.rs @@ -15,9 +15,9 @@ use std::sync::{Arc, RwLock}; use ffai_core::{Backend, Binding, Device, DeviceBuffer, Error, Grid, Kernel, Result}; use ffai_kernels_codegen::{CodegenBackend, CudaGenerator}; use ffai_kernels_core::ir::KernelMode; -use ffai_kernels_runtime::{CudaDevice as MtCudaDevice, CudaModule, MetalTileError}; +use ffai_kernels_runtime::{CudaDevice as MtCudaDevice, CudaModule, FFAIError}; -fn dispatch_err(e: MetalTileError) -> Error { +fn dispatch_err(e: FFAIError) -> Error { Error::Dispatch(e.to_string()) } diff --git a/rust/crates/backends/ffai-cuda/tests/cuda_smoke.rs b/rust/crates/backends/ffai-cuda/tests/cuda_smoke.rs index f5bd631a..c68a7775 100644 --- a/rust/crates/backends/ffai-cuda/tests/cuda_smoke.rs +++ b/rust/crates/backends/ffai-cuda/tests/cuda_smoke.rs @@ -143,8 +143,8 @@ fn ffai_ops_elementwise_on_cuda() { eprintln!("ffai_ops::add + ffai_ops::mul on CUDA: OK (max|Δ|={err:.1e})"); } -/// Heavier ops via the registered-kernel lookup: rms_norm (mt_rms_norm) and -/// gemv (mt_gemv), driven through the shared Device trait on CUDA and +/// Heavier ops via the registered-kernel lookup: rms_norm (ffai_rms_norm) and +/// gemv (ffai_gemv), driven through the shared Device trait on CUDA and /// checked against a CPU reference. This is the mechanism every transformer /// op rides on. #[test] @@ -180,7 +180,7 @@ fn ffai_ops_rms_norm_and_gemv_on_cuda() { } } assert!(rms_err <= 1e-4, "rms_norm on CUDA mismatch: max|Δ|={rms_err:.3e}"); - eprintln!("ffai_ops::rms_norm (mt_rms_norm) on CUDA: OK (max|Δ|={rms_err:.1e})"); + eprintln!("ffai_ops::rms_norm (ffai_rms_norm) on CUDA: OK (max|Δ|={rms_err:.1e})"); // ── gemv: [M,K] @ [K] ──────────────────────────────────────────── const M: usize = 64; @@ -202,7 +202,7 @@ fn ffai_ops_rms_norm_and_gemv_on_cuda() { gemv_err = gemv_err.max((got[r] - want).abs()); } assert!(gemv_err <= 1e-3, "gemv on CUDA mismatch: max|Δ|={gemv_err:.3e}"); - eprintln!("ffai_ops::gemv (mt_gemv) on CUDA: OK (max|Δ|={gemv_err:.1e})"); + eprintln!("ffai_ops::gemv (ffai_gemv) on CUDA: OK (max|Δ|={gemv_err:.1e})"); } fn sigmoid(x: f32) -> f32 { @@ -231,7 +231,7 @@ fn ffai_ops_transformer_ops_on_cuda() { e = e.max((s[i] - g[i] * sigmoid(g[i])).abs()); } assert!(e <= 1e-5, "silu mismatch: {e:.2e}"); - eprintln!("ffai_ops::silu (mt_silu) on CUDA: OK (max|Δ|={e:.1e})"); + eprintln!("ffai_ops::silu (ffai_silu) on CUDA: OK (max|Δ|={e:.1e})"); // ── swiglu ─────────────────────────────────────────────────────── let up: Vec = (0..1024).map(|i| (i % 13) as f32 * 0.05).collect(); @@ -246,7 +246,7 @@ fn ffai_ops_transformer_ops_on_cuda() { e = e.max((w[i] - g[i] * sigmoid(g[i]) * up[i]).abs()); } assert!(e <= 1e-5, "swiglu mismatch: {e:.2e}"); - eprintln!("ffai_ops::swiglu (mt_swiglu) on CUDA: OK (max|Δ|={e:.1e})"); + eprintln!("ffai_ops::swiglu (ffai_swiglu) on CUDA: OK (max|Δ|={e:.1e})"); // ── gather (embedding) ─────────────────────────────────────────── const VOCAB: usize = 8; @@ -291,7 +291,7 @@ fn ffai_ops_transformer_ops_on_cuda() { } } assert!(e <= 1e-5, "softmax mismatch: {e:.2e}"); - eprintln!("ffai_ops::softmax (mt_softmax) on CUDA: OK (max|Δ|={e:.1e})"); + eprintln!("ffai_ops::softmax (ffai_softmax) on CUDA: OK (max|Δ|={e:.1e})"); } /// Decode-time attention (sdpa_decode) on CUDA vs a CPU reference, single diff --git a/rust/crates/backends/ffai-cuda/tests/ssm_test.rs b/rust/crates/backends/ffai-cuda/tests/ssm_test.rs index 362c0ea8..6c27d5de 100644 --- a/rust/crates/backends/ffai-cuda/tests/ssm_test.rs +++ b/rust/crates/backends/ffai-cuda/tests/ssm_test.rs @@ -1,7 +1,7 @@ #![cfg(feature = "cuda")] // Copyright 2026 Eric Kryski (@ekryski) and Tom Turney (@TheTom) // SPDX-License-Identifier: Apache-2.0 -//! Mamba2 SSD selective-scan decode step (mt_ssm_step) on CUDA vs CPU — +//! Mamba2 SSD selective-scan decode step (ffai_ssm_step) on CUDA vs CPU — //! the core SSM-family op (Mamba2/Jamba/FalconH1/LFM2). use ffai_core::{DType, Device, Tensor}; use ffai_cuda::CudaDevice; @@ -13,7 +13,7 @@ fn fill(n:usize,s:usize)->Vec{(0..n).map(|i|(((i*7+s*131)%89) as f32-44.0)* fn tn(d:&dyn Device,v:&[f32],sh:Vec)->Tensor{Tensor::new(d.upload(&tb(v)).unwrap(),sh,DType::F32)} #[test] -fn mt_ssm_step_on_cuda_matches_cpu(){ +fn ffai_ssm_step_on_cuda_matches_cpu(){ let Some(dev)=CudaDevice::create().expect("metal") else { eprintln!("no CUDA — skip"); return; }; let (nh, dh, ds, hpg) = (4usize, 8usize, 32usize, 2usize); let ng = nh/hpg; @@ -43,7 +43,7 @@ fn mt_ssm_step_on_cuda_matches_cpu(){ } let mut es=0.0f32; for i in 0..nh*dh*ds { es=es.max((so[i]-so_ref[i]).abs()); } let mut eo=0.0f32; for i in 0..nh*dh { eo=eo.max((out[i]-out_ref[i]).abs()); } - eprintln!("mt_ssm_step on CUDA vs CPU: state max|Δ|={es:.3e} out max|Δ|={eo:.3e}"); + eprintln!("ffai_ssm_step on CUDA vs CPU: state max|Δ|={es:.3e} out max|Δ|={eo:.3e}"); assert!(es<=1e-4 && eo<=1e-4, "ssm mismatch state={es:.3e} out={eo:.3e}"); eprintln!("✅ Mamba2 SSD selective-scan step runs on CUDA through the shared op layer, matches CPU."); } diff --git a/rust/crates/backends/ffai-metal/src/lib.rs b/rust/crates/backends/ffai-metal/src/lib.rs index 942d7116..45be3110 100644 --- a/rust/crates/backends/ffai-metal/src/lib.rs +++ b/rust/crates/backends/ffai-metal/src/lib.rs @@ -16,13 +16,13 @@ //! (preserves in-place reads + readback) and invalidate the stale resident copy. use ffai_core::{Backend, Binding, Device, DeviceBuffer, Error, Grid, Kernel, Result}; -use ffai_kernels_runtime::{Context, DispatchSpec, MetalTileError, ResidentBuffer}; +use ffai_kernels_runtime::{Context, DispatchSpec, FFAIError, ResidentBuffer}; use parking_lot::{Mutex, RwLock}; use std::any::Any; use std::collections::BTreeMap; use std::sync::Arc; -fn err(e: MetalTileError) -> Error { +fn err(e: FFAIError) -> Error { Error::Dispatch(e.to_string()) } diff --git a/rust/crates/backends/ffai-metal/tests/gemv_q8.rs b/rust/crates/backends/ffai-metal/tests/gemv_q8.rs index 8497ed78..210175b0 100644 --- a/rust/crates/backends/ffai-metal/tests/gemv_q8.rs +++ b/rust/crates/backends/ffai-metal/tests/gemv_q8.rs @@ -11,7 +11,7 @@ fn tb_f32(v: &[f32]) -> Vec { v.iter().flat_map(|x| x.to_le_bytes()).collect fn tb_u32(v: &[u32]) -> Vec { v.iter().flat_map(|x| x.to_le_bytes()).collect() } fn fb(b: &[u8]) -> Vec { b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect() } // f32<->f16 for the q4 coalesced kernel, whose scales are f16 (see ffai-kernels -// gemv_quantized: mt_gemv_q4_coalesced reads f16 scales). Feed f16 + reference +// gemv_quantized: ffai_gemv_q4_coalesced reads f16 scales). Feed f16 + reference // against the f16-rounded scale so kernel and CPU oracle agree exactly. fn f32_to_f16(f: f32) -> u16 { let x = f.to_bits(); diff --git a/rust/crates/backends/ffai-metal/tests/ssm_test.rs b/rust/crates/backends/ffai-metal/tests/ssm_test.rs index 1bb58d08..421e8584 100644 --- a/rust/crates/backends/ffai-metal/tests/ssm_test.rs +++ b/rust/crates/backends/ffai-metal/tests/ssm_test.rs @@ -1,6 +1,6 @@ // Copyright 2026 Eric Kryski (@ekryski) and Tom Turney (@TheTom) // SPDX-License-Identifier: Apache-2.0 -//! Mamba2 SSD selective-scan decode step (mt_ssm_step) on Metal vs CPU — +//! Mamba2 SSD selective-scan decode step (ffai_ssm_step) on Metal vs CPU — //! the core SSM-family op (Mamba2/Jamba/FalconH1/LFM2). use ffai_core::{DType, Device, Tensor}; use ffai_metal::MetalDevice; @@ -12,7 +12,7 @@ fn fill(n:usize,s:usize)->Vec{(0..n).map(|i|(((i*7+s*131)%89) as f32-44.0)* fn tn(d:&dyn Device,v:&[f32],sh:Vec)->Tensor{Tensor::new(d.upload(&tb(v)).unwrap(),sh,DType::F32)} #[test] -fn mt_ssm_step_on_metal_matches_cpu(){ +fn ffai_ssm_step_on_metal_matches_cpu(){ let Some(dev)=MetalDevice::create().expect("metal") else { eprintln!("no Metal — skip"); return; }; let (nh, dh, ds, hpg) = (4usize, 8usize, 32usize, 2usize); let ng = nh/hpg; @@ -42,7 +42,7 @@ fn mt_ssm_step_on_metal_matches_cpu(){ } let mut es=0.0f32; for i in 0..nh*dh*ds { es=es.max((so[i]-so_ref[i]).abs()); } let mut eo=0.0f32; for i in 0..nh*dh { eo=eo.max((out[i]-out_ref[i]).abs()); } - eprintln!("mt_ssm_step on Metal vs CPU: state max|Δ|={es:.3e} out max|Δ|={eo:.3e}"); + eprintln!("ffai_ssm_step on Metal vs CPU: state max|Δ|={es:.3e} out max|Δ|={eo:.3e}"); assert!(es<=1e-4 && eo<=1e-4, "ssm mismatch state={es:.3e} out={eo:.3e}"); eprintln!("✅ Mamba2 SSD selective-scan step runs on Apple GPU through the shared op layer, matches CPU."); } diff --git a/rust/crates/backends/ffai-vulkan/src/imp.rs b/rust/crates/backends/ffai-vulkan/src/imp.rs index f02e2806..47e8406b 100644 --- a/rust/crates/backends/ffai-vulkan/src/imp.rs +++ b/rust/crates/backends/ffai-vulkan/src/imp.rs @@ -52,10 +52,10 @@ use std::sync::{Arc, Mutex}; use ffai_core::{Backend, Binding, Device, DeviceBuffer, Error, Grid, Kernel, Result}; use ffai_kernels_core::ir::ParamKind; use ffai_kernels_runtime::{ - BatchDispatch, MetalTileError, VulkanDevice as MtVulkanDevice, VulkanPipeline, VulkanRawBuffer, + BatchDispatch, FFAIError, VulkanDevice as MtVulkanDevice, VulkanPipeline, VulkanRawBuffer, }; -fn dispatch_err(e: MetalTileError) -> Error { +fn dispatch_err(e: FFAIError) -> Error { Error::Dispatch(e.to_string()) } diff --git a/rust/crates/ffai-modeltests/src/lib.rs b/rust/crates/ffai-modeltests/src/lib.rs index 5db1369e..0b5bb3bc 100644 --- a/rust/crates/ffai-modeltests/src/lib.rs +++ b/rust/crates/ffai-modeltests/src/lib.rs @@ -2868,13 +2868,13 @@ pub fn bench_nemotron(d: &dyn Device, plat: &str) { let amax_dn = Tensor::new(d.alloc_zeroed(4).unwrap(), vec![1], DType::U32); let (off0, st0, sw0) = moe_dev.as_ref().unwrap(); let (st_e, sw_e, off_e) = ffai_ops::moe_extend_groups(d, st0, sw0, off0, mt, s, n_exp).unwrap(); - let mt_e = mt + 2 * s; - let up_out = pf!(11, 2.0*mt_e as f64*inter as f64*hid as f64, - ffai_ops::moe_fp4_grouped_mma_dev(d, &xs_in, uwp, uwsc, ugw, &off_e, n_exp + 2, mt_e, inter, hid, 0, 0.0, None, Some((&amax_dn, 1, inv)), Some(&st_e)).unwrap()); - let dn_out = pf!(11, 2.0*mt_e as f64*hid as f64*inter as f64, - ffai_ops::moe_fp4_grouped_mma_dev(d, &up_out, dwp, dwsc, dgw, &off_e, n_exp + 2, mt_e, hid, inter, 1, inv, Some(&amax_dn), None, None).unwrap()); + let ffai_e = mt + 2 * s; + let up_out = pf!(11, 2.0*ffai_e as f64*inter as f64*hid as f64, + ffai_ops::moe_fp4_grouped_mma_dev(d, &xs_in, uwp, uwsc, ugw, &off_e, n_exp + 2, ffai_e, inter, hid, 0, 0.0, None, Some((&amax_dn, 1, inv)), Some(&st_e)).unwrap()); + let dn_out = pf!(11, 2.0*ffai_e as f64*hid as f64*inter as f64, + ffai_ops::moe_fp4_grouped_mma_dev(d, &up_out, dwp, dwsc, dgw, &off_e, n_exp + 2, ffai_e, hid, inter, 1, inv, Some(&amax_dn), None, None).unwrap()); let acc_dev = Tensor::new(d.alloc_zeroed(s * hid * 4).unwrap(), vec![s, hid], DType::F32); - moe_scatter_add_det_dev(d, &dn_out, &st_e, &sw_e, &acc_dev, s, mt_e, hid, 256.0f32, true).unwrap(); + moe_scatter_add_det_dev(d, &dn_out, &st_e, &sw_e, &acc_dev, s, ffai_e, hid, 256.0f32, true).unwrap(); shared_folded = true; if ondevice_moe { acc_dev_keep = Some(acc_dev); } else { acc_h = dl(&acc_dev, s * hid); } } else if fp4_moe { diff --git a/rust/crates/ffai-ops/src/lib.rs b/rust/crates/ffai-ops/src/lib.rs index d07b815b..925c89de 100644 --- a/rust/crates/ffai-ops/src/lib.rs +++ b/rust/crates/ffai-ops/src/lib.rs @@ -225,7 +225,7 @@ pub fn swiglu_limit(dev: &dyn Device, gate: &Tensor, up: &Tensor, limit: f32) -> if gate.shape != up.shape { return Err(Error::Msg("swiglu_limit: gate/up shape mismatch".into())); } - let k = lookup("mt_swiglu_limit", gate.dtype)?; + let k = lookup("ffai_swiglu_limit", gate.dtype)?; let out = Tensor::empty(dev, gate.shape.clone(), gate.dtype)?; let n = gate.elem_count() as u32; let grid = Grid::d1(n.div_ceil(256), 256); @@ -273,7 +273,7 @@ pub fn conv1d_causal_step( n_channels: u32, kernel_size: u32, ) -> Result { - let k = lookup("mt_conv1d_causal_step", x.dtype)?; + let k = lookup("ffai_conv1d_causal_step", x.dtype)?; let y = Tensor::empty(dev, vec![n_channels as usize], x.dtype)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); let bindings = vec![ @@ -310,7 +310,7 @@ fn gcd_u32(mut a: u32, mut b: u32) -> u32 { /// Mamba2 SSD selective-scan **decode step** (single token, batch=1): /// `da = exp(-exp(a_log)·dt)`; `state' = da·state + x·dt·B`; -/// `out = Σ_s C·state' + x·D`. Dispatches `mt_ssm_step`. Returns +/// `out = Σ_s C·state' + x·D`. Dispatches `ffai_ssm_step`. Returns /// `(state_out [n_heads·dh·ds], out [n_heads·dh])`. `ds` must be a multiple /// of 32. Shapes: x `[n_heads·dh]`, a_log/d_skip/dt `[n_heads]`, /// b_mat/c_mat `[n_groups·ds]`, state_in `[n_heads·dh·ds]`. @@ -329,7 +329,7 @@ pub fn ssm_step( n_heads: u32, heads_per_group: u32, ) -> Result<(Tensor, Tensor)> { - let k = lookup("mt_ssm_step_grouped", x.dtype)?; + let k = lookup("ffai_ssm_step_grouped", x.dtype)?; let state_out = Tensor::empty(dev, state_in.shape.clone(), x.dtype)?; let out = Tensor::empty(dev, vec![(n_heads * dh) as usize], x.dtype)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); @@ -413,7 +413,7 @@ pub fn dsv4_mhc_collapse( hidden_dim: u32, n_hc: u32, ) -> Result { - let k = lookup("mt_mhc_collapse", state.dtype)?; + let k = lookup("ffai_mhc_collapse", state.dtype)?; let out = Tensor::empty(dev, vec![hidden_dim as usize], state.dtype)?; let u = |x: u32| Binding::Scalar(x.to_le_bytes().to_vec()); let grid = Grid::d1((hidden_dim).div_ceil(256), 256); @@ -445,7 +445,7 @@ pub fn dsv4_mhc_expand( hidden_dim: u32, n_hc: u32, ) -> Result { - let k = lookup("mt_mhc_expand", block_out.dtype)?; + let k = lookup("ffai_mhc_expand", block_out.dtype)?; let state = Tensor::empty(dev, vec![(n_hc * hidden_dim) as usize], block_out.dtype)?; let u = |x: u32| Binding::Scalar(x.to_le_bytes().to_vec()); let grid = Grid::d1((hidden_dim).div_ceil(256), 256); @@ -469,10 +469,10 @@ pub fn dsv4_mhc_expand( // ── Heavier ops — dispatch the registered metaltile kernels ───────────── /// Row-wise RMS norm: `out[r] = x[r] * rsqrt(mean(x[r]²) + eps) * weight`. -/// Dispatches the registered `mt_rms_norm` reduction kernel — the same one +/// Dispatches the registered `ffai_rms_norm` reduction kernel — the same one /// the Swift side runs. The last dim is the row width `n`; the kernel owns 4 /// elements per thread, so `n` must be a multiple of 128 and ≤ 4096 (the -/// `mt_rms_norm_wide` variant lifts this — wired later). +/// `ffai_rms_norm_wide` variant lifts this — wired later). pub fn rms_norm(dev: &dyn Device, x: &Tensor, weight: &Tensor, eps: f32) -> Result { let n = *x.shape.last().ok_or_else(|| Error::Msg("rms_norm: scalar input".into()))?; let rows = x.elem_count() / n; @@ -482,9 +482,9 @@ pub fn rms_norm(dev: &dyn Device, x: &Tensor, weight: &Tensor, eps: f32) -> Resu // Fast path: 4 elems/thread, needs n a multiple of 128 and ≤ 4096. // Otherwise the strided wide variant handles any row width. let (kname, block) = if n % 128 == 0 && n <= 4096 { - ("mt_rms_norm", (n / 4) as u32) + ("ffai_rms_norm", (n / 4) as u32) } else { - ("mt_rms_norm_wide", 256u32) + ("ffai_rms_norm_wide", 256u32) }; let k = lookup(kname, x.dtype)?; let grid = Grid { grid: [rows as u32, 1, 1], block: [block, 1, 1] }; @@ -503,7 +503,7 @@ pub fn rms_norm(dev: &dyn Device, x: &Tensor, weight: &Tensor, eps: f32) -> Resu } /// Matrix-vector product `mat @ vec`: `mat` is `[m, k]` row-major, `vec` is -/// `[k]`, result is `[m]`. Dispatches the registered `mt_gemv` kernel (one +/// `[k]`, result is `[m]`. Dispatches the registered `ffai_gemv` kernel (one /// threadgroup per output row). This is the decode-time projection path; the /// batched/prefill cooperative matmul is a separate kernel, wired later. pub fn gemv(dev: &dyn Device, mat: &Tensor, vec: &Tensor) -> Result { @@ -517,7 +517,7 @@ pub fn gemv(dev: &dyn Device, mat: &Tensor, vec: &Tensor) -> Result { vec.elem_count() ))); } - let k = lookup("mt_gemv", mat.dtype)?; + let k = lookup("ffai_gemv", mat.dtype)?; let out = Tensor::empty(dev, vec![m], mat.dtype)?; let grid = Grid { grid: [m as u32, 1, 1], block: [256, 1, 1] }; @@ -557,7 +557,7 @@ pub fn gemv_q8( // are JIT-specialized by the runtime from the scalar bindings below). // Coalesced variant: consecutive lanes read consecutive qs words (~2× the // strided original's DRAM bandwidth on GB10). - let k = cached_ir("ffai_gemv_q8_coalesced", x.dtype, || { let mut k = ffai_kernels_std::kernels::gemm::gemv_quantized::mt_gemv_q8_coalesced::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); + let k = cached_ir("ffai_gemv_q8_coalesced", x.dtype, || { let mut k = ffai_kernels_std::kernels::gemm::gemv_quantized::ffai_gemv_q8_coalesced::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); let out = Tensor::empty(dev, vec![m_out], x.dtype)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); let grid = Grid { grid: [m_out as u32, 1, 1], block: [32, 1, 1] }; @@ -608,7 +608,7 @@ pub fn gemm_q8_mpp( // `ffai_gemm_q8_mpp` in Reduction mode (matches the metaltile coopmat test // + bench). Buffer order: x, qs, d_f32, out; constexpr n_rows, out_dim, k_in. let k = cached_ir("ffai_gemm_q8_mpp", DType::F16, || { - let mut k = ffai_kernels_std::kernels::gemm::gemm_q8_mpp::mt_gemm_q8_mpp::kernel_ir_for(DType::F16); + let mut k = ffai_kernels_std::kernels::gemm::gemm_q8_mpp::ffai_gemm_q8_mpp::kernel_ir_for(DType::F16); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); @@ -727,7 +727,7 @@ pub fn add_bias_rows(dev: &dyn Device, x: &Tensor, bias: &Tensor, n_rows: usize, /// Q8 gemv with fused ReLU²: `out[r] = max(0, (Wq·x)[r])²` — a MoE expert's /// `up` projection + activation in one dispatch. Dispatches `ffai_gemv_q8_coalesced_relu2`. pub fn gemv_q8_relu2(dev: &dyn Device, qs: &Tensor, scales: &Tensor, x: &Tensor, m_out: usize, k_in: usize, rows_per_group: usize) -> Result { - let k = cached_ir("ffai_gemv_q8_coalesced_relu2", x.dtype, || { let mut k = ffai_kernels_std::kernels::gemm::gemv_quantized::mt_gemv_q8_coalesced_relu2::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); + let k = cached_ir("ffai_gemv_q8_coalesced_relu2", x.dtype, || { let mut k = ffai_kernels_std::kernels::gemm::gemv_quantized::ffai_gemv_q8_coalesced_relu2::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); let out = Tensor::empty(dev, vec![m_out], x.dtype)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); let grid = Grid { grid: [m_out as u32, 1, 1], block: [32, 1, 1] }; @@ -754,7 +754,7 @@ pub fn gemv_q8_accum( if k_in % 32 != 0 { return Err(Error::Msg(format!("gemv_q8_accum: k_in {k_in} must be a multiple of 32"))); } - let k = cached_ir("ffai_gemv_q8_coalesced_accum", acc.dtype, || { let mut k = ffai_kernels_std::kernels::gemm::gemv_quantized::mt_gemv_q8_coalesced_accum::kernel_ir_for(acc.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); + let k = cached_ir("ffai_gemv_q8_coalesced_accum", acc.dtype, || { let mut k = ffai_kernels_std::kernels::gemm::gemv_quantized::ffai_gemv_q8_coalesced_accum::kernel_ir_for(acc.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); let grid = Grid { grid: [m_out as u32, 1, 1], block: [32, 1, 1] }; dev.dispatch( @@ -780,7 +780,7 @@ pub fn gemv_q8_accum( /// context entirely on-device — no host reorg/reupload per step (the 32K-context /// fix). Dispatches `ffai_kv_append` (runtime `pos` ⇒ compiled once, not per step). pub fn kv_append(dev: &dyn Device, src: &Tensor, dst: &Tensor, posbuf: &Tensor, hd: usize, cap: usize, n: usize) -> Result<()> { - let k = cached_ir("ffai_kv_append", src.dtype, || { let mut k = ffai_kernels_std::kernels::kv_cache::cache::mt_kv_append::kernel_ir_for(src.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); + let k = cached_ir("ffai_kv_append", src.dtype, || { let mut k = ffai_kernels_std::kernels::kv_cache::cache::ffai_kv_append::kernel_ir_for(src.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); let grid = Grid { grid: [(n as u32).div_ceil(64), 1, 1], block: [64, 1, 1] }; dev.dispatch(&k, &[Binding::Buffer(src.buffer.clone()), Binding::Buffer(dst.buffer.clone()), Binding::Buffer(posbuf.buffer.clone()), u(hd as u32), u(cap as u32)], grid)?; @@ -789,7 +789,7 @@ pub fn kv_append(dev: &dyn Device, src: &Tensor, dst: &Tensor, posbuf: &Tensor, /// Device slice `out[i] = src[off + i]` for `len` elements (no host round-trip). pub fn slice(dev: &dyn Device, src: &Tensor, off: usize, len: usize) -> Result { - let k = cached_ir("ffai_slice", src.dtype, || { let mut k = ffai_kernels_std::kernels::ops::copy::mt_slice::kernel_ir_for(src.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); + let k = cached_ir("ffai_slice", src.dtype, || { let mut k = ffai_kernels_std::kernels::ops::copy::ffai_slice::kernel_ir_for(src.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); // Defensive bounds check. The kernel reads `src[off + i]` for i in [0, len), // addressing a sub-slab purely via `off` against the RAW bound buffer (the // tensor.offset is ignored by design). A caller binding a buffer that does @@ -813,7 +813,7 @@ pub fn slice(dev: &dyn Device, src: &Tensor, off: usize, len: usize) -> Result Result { - let k = cached_ir("ffai_cast_f32_f16", DType::F16, || { let mut k = ffai_kernels_std::kernels::ops::unary::mt_cast_f32_f16::kernel_ir(); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); + let k = cached_ir("ffai_cast_f32_f16", DType::F16, || { let mut k = ffai_kernels_std::kernels::ops::unary::ffai_cast_f32_f16::kernel_ir(); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); let n = src.elem_count(); let out = Tensor::empty(dev, src.shape.clone(), DType::F16)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); @@ -824,7 +824,7 @@ pub fn cast_f32_f16(dev: &dyn Device, src: &Tensor) -> Result { /// Cast f16 → f32 (reverse: widen the sdpa f16 output back to f32 for the /// downstream o_proj Q4 GEMV, which consumes f32 activations). Fresh F32 tensor. pub fn cast_f16_f32(dev: &dyn Device, src: &Tensor) -> Result { - let k = cached_ir("ffai_cast_f16_f32", DType::F32, || { let mut k = ffai_kernels_std::kernels::ops::unary::mt_cast_f16_f32::kernel_ir(); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); + let k = cached_ir("ffai_cast_f16_f32", DType::F32, || { let mut k = ffai_kernels_std::kernels::ops::unary::ffai_cast_f16_f32::kernel_ir(); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); let n = src.elem_count(); let out = Tensor::empty(dev, src.shape.clone(), DType::F32)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); @@ -836,7 +836,7 @@ pub fn cast_f16_f32(dev: &dyn Device, src: &Tensor) -> Result { /// (NEMOTRON_BF16_STREAM) doesn't overflow on NemotronH massive-activation /// channels the way f16 (max 65504) does. pub fn cast_f32_bf16(dev: &dyn Device, src: &Tensor) -> Result { - let k = cached_ir("ffai_cast_f32_bf16", DType::BF16, || { let mut k = ffai_kernels_std::kernels::ops::unary::mt_cast_f32_bf16::kernel_ir(); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); + let k = cached_ir("ffai_cast_f32_bf16", DType::BF16, || { let mut k = ffai_kernels_std::kernels::ops::unary::ffai_cast_f32_bf16::kernel_ir(); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); let n = src.elem_count(); let out = Tensor::empty(dev, src.shape.clone(), DType::BF16)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); @@ -846,7 +846,7 @@ pub fn cast_f32_bf16(dev: &dyn Device, src: &Tensor) -> Result { /// Cast bf16 → f32 (widen the residual back for ops that need f32 input). pub fn cast_bf16_f32(dev: &dyn Device, src: &Tensor) -> Result { - let k = cached_ir("ffai_cast_bf16_f32", DType::F32, || { let mut k = ffai_kernels_std::kernels::ops::unary::mt_cast_to_f32::kernel_ir_for(DType::BF16); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); + let k = cached_ir("ffai_cast_bf16_f32", DType::F32, || { let mut k = ffai_kernels_std::kernels::ops::unary::ffai_cast_to_f32::kernel_ir_for(DType::BF16); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); let n = src.elem_count(); let out = Tensor::empty(dev, src.shape.clone(), DType::F32)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); @@ -856,7 +856,7 @@ pub fn cast_bf16_f32(dev: &dyn Device, src: &Tensor) -> Result { /// Device Mamba dt: `softplus(dt_raw + dt_bias)` — no host round-trip. pub fn softplus_add(dev: &dyn Device, a: &Tensor, b: &Tensor) -> Result { - let k = cached_ir("ffai_softplus_add", DType::F32, || { let mut k = ffai_kernels_std::kernels::ssm::scan::mt_softplus_add::kernel_ir(); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); + let k = cached_ir("ffai_softplus_add", DType::F32, || { let mut k = ffai_kernels_std::kernels::ssm::scan::ffai_softplus_add::kernel_ir(); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); let n = a.elem_count(); let out = Tensor::empty(dev, a.shape.clone(), DType::F32)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); @@ -867,7 +867,7 @@ pub fn softplus_add(dev: &dyn Device, a: &Tensor, b: &Tensor) -> Result /// NemotronH/Zamba2 gated grouped RMSNorm ON-DEVICE: out = (y·silu(z)) normalized /// per `gs`-group, ×w. Removes the per-Mamba-layer dl→host-norm→up sync. `y` is f32. pub fn gated_group_rmsnorm(dev: &dyn Device, y: &Tensor, z: &Tensor, w: &Tensor, eps: f32, di: usize, gs: usize) -> Result { - let k = cached_ir("ffai_gated_group_rmsnorm", z.dtype, || { let mut k = ffai_kernels_std::kernels::ssm::scan::mt_gated_group_rmsnorm::kernel_ir_for(z.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); + let k = cached_ir("ffai_gated_group_rmsnorm", z.dtype, || { let mut k = ffai_kernels_std::kernels::ssm::scan::ffai_gated_group_rmsnorm::kernel_ir_for(z.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); let out = Tensor::empty(dev, vec![di], z.dtype)?; let eps_buf = Tensor::new(scalar_buf(dev, eps.to_bits())?, vec![1], DType::F32); let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); @@ -880,7 +880,7 @@ pub fn gated_group_rmsnorm(dev: &dyn Device, y: &Tensor, z: &Tensor, w: &Tensor, pub fn conv_roll(dev: &dyn Device, old: &Tensor, xbc: &Tensor, conv_dim: usize, kc: usize) -> Result { let n = (kc - 1) * conv_dim; let keep = (kc - 2) * conv_dim; - let k = cached_ir("ffai_conv_roll", old.dtype, || { let mut k = ffai_kernels_std::kernels::convolution::conv1d_causal::mt_conv_roll::kernel_ir_for(old.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); + let k = cached_ir("ffai_conv_roll", old.dtype, || { let mut k = ffai_kernels_std::kernels::convolution::conv1d_causal::ffai_conv_roll::kernel_ir_for(old.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); let out = Tensor::empty(dev, vec![n], old.dtype)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); dev.dispatch(&k, &[Binding::Buffer(old.buffer.clone()), Binding::Buffer(xbc.buffer.clone()), Binding::Buffer(out.buffer.clone()), u(conv_dim as u32), u(keep as u32), u(n as u32)], Grid { grid: [(n as u32).div_ceil(256), 1, 1], block: [256, 1, 1] })?; @@ -907,7 +907,7 @@ pub fn conv1d_causal_prefill( conv_dim: usize, kc: usize, ) -> Result { - let kernel = lookup("mt_conv1d_causal_prefill", DType::F32)?; + let kernel = lookup("ffai_conv1d_causal_prefill", DType::F32)?; let y = Tensor::empty(dev, vec![s * conv_dim], DType::F32)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); dev.dispatch( @@ -941,7 +941,7 @@ pub fn strided_col_copy( col_off: usize, width: usize, ) -> Result { - let kernel = lookup("mt_strided_col_copy", DType::F32)?; + let kernel = lookup("ffai_strided_col_copy", DType::F32)?; // Defensive bounds check. The kernel reads `src[ti*stride + col_off + ci]` // for ti in [0, s), ci in [0, width) — a sub-slab addressed purely via the // host stride/col_off against the RAW bound buffer (tensor.offset ignored by @@ -986,7 +986,7 @@ pub fn softplus_add_rows( s: usize, n: usize, ) -> Result { - let kernel = lookup("mt_softplus_add_rows", DType::F32)?; + let kernel = lookup("ffai_softplus_add_rows", DType::F32)?; let dst = Tensor::empty(dev, vec![s * n], DType::F32)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); dev.dispatch( @@ -1012,7 +1012,7 @@ pub fn softplus_add_rows( /// Batched MoE up+ReLU²: gather `top_k` experts (indices `idx`) from the /// contiguous `[n_exp*inter, hid]` Q4 weight into one `[top_k*inter]` GEMV. pub fn moe_gather_up_relu2(dev: &dyn Device, qs: &Tensor, scales: &Tensor, x: &Tensor, idx: &Tensor, top_k: usize, inter: usize, hid: usize) -> Result { - let k = cached_ir("ffai_moe_gather_q4_relu2", x.dtype, || { let mut k = ffai_kernels_std::kernels::moe::moe_gather_q4::mt_moe_gather_q4_relu2::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); + let k = cached_ir("ffai_moe_gather_q4_relu2", x.dtype, || { let mut k = ffai_kernels_std::kernels::moe::moe_gather_q4::ffai_moe_gather_q4_relu2::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); let out = Tensor::empty(dev, vec![top_k * inter], x.dtype)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); // Default rpt=2: MoE-gather kernels are big + latency-bound; 2 warps/row hides @@ -1025,7 +1025,7 @@ pub fn moe_gather_up_relu2(dev: &dyn Device, qs: &Tensor, scales: &Tensor, x: &T /// contiguous `[n_exp*hid, inter]` Q4 weight; `x` is the `[top_k*inter]` up output. #[allow(clippy::too_many_arguments)] pub fn moe_gather_down_accum(dev: &dyn Device, qs: &Tensor, scales: &Tensor, x: &Tensor, idx: &Tensor, wts: &Tensor, acc: &Tensor, top_k: usize, inter: usize, hid: usize) -> Result<()> { - let k = cached_ir("ffai_moe_gather_q4_down_accum", x.dtype, || { let mut k = ffai_kernels_std::kernels::moe::moe_gather_q4::mt_moe_gather_q4_down_accum::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); + let k = cached_ir("ffai_moe_gather_q4_down_accum", x.dtype, || { let mut k = ffai_kernels_std::kernels::moe::moe_gather_q4::ffai_moe_gather_q4_down_accum::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); dev.dispatch(&k, &[Binding::Buffer(qs.buffer.clone()), Binding::Buffer(scales.buffer.clone()), Binding::Buffer(x.buffer.clone()), Binding::Buffer(idx.buffer.clone()), Binding::Buffer(wts.buffer.clone()), Binding::Buffer(acc.buffer.clone()), u(inter as u32), u(hid as u32), u(top_k as u32)], Grid { grid: [hid as u32, 1, 1], block: [32, 1, 1] })?; Ok(()) @@ -1033,7 +1033,7 @@ pub fn moe_gather_down_accum(dev: &dyn Device, qs: &Tensor, scales: &Tensor, x: /// Batched MoE down gather → `[top_k*hid]` (one big GEMV, no accumulate). pub fn moe_gather_down(dev: &dyn Device, qs: &Tensor, scales: &Tensor, x: &Tensor, idx: &Tensor, top_k: usize, inter: usize, hid: usize) -> Result { - let k = cached_ir("ffai_moe_gather_q4_down", x.dtype, || { let mut k = ffai_kernels_std::kernels::moe::moe_gather_q4::mt_moe_gather_q4_down::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); + let k = cached_ir("ffai_moe_gather_q4_down", x.dtype, || { let mut k = ffai_kernels_std::kernels::moe::moe_gather_q4::ffai_moe_gather_q4_down::kernel_ir_for(x.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k }); let out = Tensor::empty(dev, vec![top_k * hid], x.dtype)?; let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); // Default rpt=2: MoE-gather kernels are big + latency-bound; 2 warps/row hides @@ -1240,13 +1240,13 @@ pub fn moe_router_device(dev: &dyn Device, gate_logits: &Tensor, bias: &Tensor, let f = |v: f32| Binding::Scalar(v.to_le_bytes().to_vec()); let unbiased = Tensor::empty(dev, vec![n_exp], DType::F32)?; let biased = Tensor::empty(dev, vec![n_exp], DType::F32)?; - let ks = cached_ir("ffai_moe_sigmoid_bias", DType::F32, || { let mut k = ffai_kernels_std::kernels::moe::moe_sigmoid_bias::mt_moe_sigmoid_bias::kernel_ir(); k.mode = KernelMode::Grid3D; k }); + let ks = cached_ir("ffai_moe_sigmoid_bias", DType::F32, || { let mut k = ffai_kernels_std::kernels::moe::moe_sigmoid_bias::ffai_moe_sigmoid_bias::kernel_ir(); k.mode = KernelMode::Grid3D; k }); dev.dispatch(&ks, &[Binding::Buffer(gate_logits.buffer.clone()), Binding::Buffer(bias.buffer.clone()), Binding::Buffer(unbiased.buffer.clone()), Binding::Buffer(biased.buffer.clone()), u(n_exp as u32)], Grid { grid: [(n_exp as u32).div_ceil(256), 1, 1], block: [256, 1, 1] })?; let idx = Tensor::empty(dev, vec![top_k], DType::U32)?; let wts = Tensor::empty(dev, vec![top_k], DType::F32)?; - let kr = cached_ir("mt_dsv4_router_topk", DType::F32, || { let mut k = ffai_kernels_std::kernels::moe::moe_router_topk_biased::mt_moe_router_topk_biased::kernel_ir_for(DType::F32); k.mode = KernelMode::Reduction; k }); + let kr = cached_ir("ffai_dsv4_router_topk", DType::F32, || { let mut k = ffai_kernels_std::kernels::moe::moe_router_topk_biased::ffai_moe_router_topk_biased::kernel_ir_for(DType::F32); k.mode = KernelMode::Reduction; k }); dev.dispatch(&kr, &[Binding::Buffer(biased.buffer.clone()), Binding::Buffer(unbiased.buffer.clone()), Binding::Buffer(idx.buffer.clone()), Binding::Buffer(wts.buffer.clone()), u(n_exp as u32), u(top_k as u32)], Grid { grid: [1, 1, 1], block: [32, 1, 1] })?; - let kv = cached_ir("ffai_vscale", DType::F32, || { let mut k = ffai_kernels_std::kernels::ops::unary::mt_vscale::kernel_ir(); k.mode = KernelMode::Grid3D; k }); + let kv = cached_ir("ffai_vscale", DType::F32, || { let mut k = ffai_kernels_std::kernels::ops::unary::ffai_vscale::kernel_ir(); k.mode = KernelMode::Grid3D; k }); dev.dispatch(&kv, &[Binding::Buffer(wts.buffer.clone()), f(scale), u(top_k as u32)], Grid { grid: [(top_k as u32).div_ceil(64), 1, 1], block: [64, 1, 1] })?; Ok((idx, wts)) } @@ -1378,7 +1378,7 @@ pub fn moe_route_sort_device( } pub fn moe_weighted_sum(dev: &dyn Device, downs: &Tensor, wts: &Tensor, acc: &Tensor, hid: usize, top_k: usize) -> Result<()> { - let k = cached_ir("ffai_moe_weighted_sum", acc.dtype, || { let mut k = ffai_kernels_std::kernels::moe::moe_gather_q4::mt_moe_weighted_sum::kernel_ir_for(acc.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); + let k = cached_ir("ffai_moe_weighted_sum", acc.dtype, || { let mut k = ffai_kernels_std::kernels::moe::moe_gather_q4::ffai_moe_weighted_sum::kernel_ir_for(acc.dtype); k.mode = ffai_kernels_core::ir::KernelMode::Grid3D; k }); let u = |v: u32| Binding::Scalar(v.to_le_bytes().to_vec()); dev.dispatch(&k, &[Binding::Buffer(downs.buffer.clone()), Binding::Buffer(wts.buffer.clone()), Binding::Buffer(acc.buffer.clone()), u(hid as u32), u(top_k as u32)], Grid { grid: [(hid as u32).div_ceil(256), 1, 1], block: [256, 1, 1] })?; Ok(()) @@ -1393,13 +1393,13 @@ fn gemv_q4_dispatch(dev: &dyn Device, kernel: &str, qs: &Tensor, scales: &Tensor let name = if vec { "ffai_gemv_q4_vec" } else if two_row { "ffai_gemv_q4_coalesced_2row" } else { match kernel { "plain" => "ffai_gemv_q4_coalesced", "relu2" => "ffai_gemv_q4_coalesced_relu2", _ => "ffai_gemv_q4_coalesced_accum" } }; let k = cached_ir(name, x.dtype, || { let mut k = if vec { - ffai_kernels_std::kernels::gemm::gemv_quantized::mt_gemv_q4_vec::kernel_ir_for(x.dtype) + ffai_kernels_std::kernels::gemm::gemv_quantized::ffai_gemv_q4_vec::kernel_ir_for(x.dtype) } else if two_row { - ffai_kernels_std::kernels::gemm::gemv_quantized::mt_gemv_q4_coalesced_2row::kernel_ir_for(x.dtype) + ffai_kernels_std::kernels::gemm::gemv_quantized::ffai_gemv_q4_coalesced_2row::kernel_ir_for(x.dtype) } else { match kernel { - "plain" => ffai_kernels_std::kernels::gemm::gemv_quantized::mt_gemv_q4_coalesced::kernel_ir_for(x.dtype), - "relu2" => ffai_kernels_std::kernels::gemm::gemv_quantized::mt_gemv_q4_coalesced_relu2::kernel_ir_for(x.dtype), - _ => ffai_kernels_std::kernels::gemm::gemv_quantized::mt_gemv_q4_coalesced_accum::kernel_ir_for(x.dtype), + "plain" => ffai_kernels_std::kernels::gemm::gemv_quantized::ffai_gemv_q4_coalesced::kernel_ir_for(x.dtype), + "relu2" => ffai_kernels_std::kernels::gemm::gemv_quantized::ffai_gemv_q4_coalesced_relu2::kernel_ir_for(x.dtype), + _ => ffai_kernels_std::kernels::gemm::gemv_quantized::ffai_gemv_q4_coalesced_accum::kernel_ir_for(x.dtype), }}; k.mode = ffai_kernels_core::ir::KernelMode::Reduction; k @@ -1639,7 +1639,7 @@ pub fn matmul(dev: &dyn Device, weight: &Tensor, input: &Tensor) -> Result) = match head_dim { - 64 => ("mt_sdpa_decode_d64", 1024, vec![u(0), f(0.0), f(scale)]), // has_sink, sink_logit, scale - 96 => ("mt_sdpa_decode_d96", 1024, vec![f(scale)]), - 128 => ("mt_sdpa_decode", 1024, vec![u(0), u(0), u(0), f(0.0), f(scale)]), // sink_end, window_start, has_sink, sink_logit, scale - 256 => ("mt_sdpa_decode_d256", 1024, vec![u(0), f(0.0), f(scale)]), - 512 => ("mt_sdpa_decode_d512", 512, vec![f(scale)]), + 64 => ("ffai_sdpa_decode_d64", 1024, vec![u(0), f(0.0), f(scale)]), // has_sink, sink_logit, scale + 96 => ("ffai_sdpa_decode_d96", 1024, vec![f(scale)]), + 128 => ("ffai_sdpa_decode", 1024, vec![u(0), u(0), u(0), f(0.0), f(scale)]), // sink_end, window_start, has_sink, sink_logit, scale + 256 => ("ffai_sdpa_decode_d256", 1024, vec![u(0), f(0.0), f(scale)]), + 512 => ("ffai_sdpa_decode_d512", 512, vec![f(scale)]), _ => return Err(Error::Msg(format!("sdpa_decode: unsupported head_dim {head_dim}"))), }; @@ -4272,7 +4272,7 @@ pub fn sdpa_decode_2pass( let f = |x: f32| Binding::Scalar(x.to_le_bytes().to_vec()); let k1 = cached_ir("sdpa_decode_2pass_pass1", q.dtype, || { - let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::mt_sdpa_decode_2pass_pass1::kernel_ir_for(q.dtype); + let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::ffai_sdpa_decode_2pass_pass1::kernel_ir_for(q.dtype); k.mode = KernelMode::Reduction; k }); @@ -4284,7 +4284,7 @@ pub fn sdpa_decode_2pass( let out = Tensor::empty(dev, vec![n_q_heads, head_dim], q.dtype)?; let k2 = cached_ir("sdpa_decode_2pass_pass2", q.dtype, || { - let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::mt_sdpa_decode_2pass_pass2::kernel_ir_for(q.dtype); + let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::ffai_sdpa_decode_2pass_pass2::kernel_ir_for(q.dtype); k.mode = KernelMode::Reduction; k }); @@ -4322,7 +4322,7 @@ pub fn sdpa_decode_2pass_tiled( let f = |x: f32| Binding::Scalar(x.to_le_bytes().to_vec()); let k1 = cached_ir("sdpa_decode_2pass_pass1_tiled", q.dtype, || { - let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::mt_sdpa_decode_2pass_pass1_tiled::kernel_ir_for(q.dtype); + let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::ffai_sdpa_decode_2pass_pass1_tiled::kernel_ir_for(q.dtype); k.mode = KernelMode::Reduction; k }); @@ -4334,7 +4334,7 @@ pub fn sdpa_decode_2pass_tiled( let out = Tensor::empty(dev, vec![n_q_heads, head_dim], q.dtype)?; let k2 = cached_ir("sdpa_decode_2pass_pass2", q.dtype, || { - let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::mt_sdpa_decode_2pass_pass2::kernel_ir_for(q.dtype); + let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::ffai_sdpa_decode_2pass_pass2::kernel_ir_for(q.dtype); k.mode = KernelMode::Reduction; k }); @@ -4374,7 +4374,7 @@ pub fn sdpa_decode_2pass_bc4( // BC=4 pass 1: 4 positions per loop iter for MLP / load-latency hiding. let k1 = cached_ir("sdpa_decode_2pass_pass1_bc4", q.dtype, || { - let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::mt_sdpa_decode_2pass_pass1_bc4::kernel_ir_for(q.dtype); + let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::ffai_sdpa_decode_2pass_pass1_bc4::kernel_ir_for(q.dtype); k.mode = KernelMode::Reduction; k }); @@ -4387,7 +4387,7 @@ pub fn sdpa_decode_2pass_bc4( // Pass 2 unchanged: same partial buffer layout, same reduction. let out = Tensor::empty(dev, vec![n_q_heads, head_dim], q.dtype)?; let k2 = cached_ir("sdpa_decode_2pass_pass2", q.dtype, || { - let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::mt_sdpa_decode_2pass_pass2::kernel_ir_for(q.dtype); + let mut k = ffai_kernels_std::kernels::sdpa::sdpa_decode_2pass::ffai_sdpa_decode_2pass_pass2::kernel_ir_for(q.dtype); k.mode = KernelMode::Reduction; k }); @@ -4418,16 +4418,16 @@ fn elementwise_kernel( Ok(out) } -/// SiLU activation `out = x * sigmoid(x)`, elementwise. Dispatches `mt_silu`. +/// SiLU activation `out = x * sigmoid(x)`, elementwise. Dispatches `ffai_silu`. pub fn silu(dev: &dyn Device, x: &Tensor) -> Result { - elementwise_kernel(dev, "mt_silu", x.dtype, x.shape.clone(), vec![Binding::Buffer(x.buffer.clone())]) + elementwise_kernel(dev, "ffai_silu", x.dtype, x.shape.clone(), vec![Binding::Buffer(x.buffer.clone())]) } /// ReLU² activation `out = max(x,0)²`, elementwise — NemotronH MoE expert act. /// Built inline (Relu → square) so the expert `up → relu² → down` chain stays -/// ON-DEVICE (no host round-trip between the two Q8 GEMVs). Dispatches `mt_relu2`. +/// ON-DEVICE (no host round-trip between the two Q8 GEMVs). Dispatches `ffai_relu2`. pub fn relu2(dev: &dyn Device, x: &Tensor) -> Result { - let mut k = Kernel::new("mt_relu2"); + let mut k = Kernel::new("ffai_relu2"); for (pname, is_out) in [("a", false), ("c", true)] { k.params.push(Param { name: pname.into(), dtype: x.dtype, shape: Shape::scalar(), is_output: is_out, kind: ParamKind::Tensor }); } @@ -4491,7 +4491,7 @@ pub fn moe_scatter_add( hid: usize, unscale: f32, ) -> Result<()> { - let mt_i = (mt as i32).to_le_bytes().to_vec(); + let ffai_i = (mt as i32).to_le_bytes().to_vec(); let hid_i = (hid as i32).to_le_bytes().to_vec(); let uscale_f = unscale.to_le_bytes().to_vec(); const BLOCK_H: u32 = 128; @@ -4507,7 +4507,7 @@ pub fn moe_scatter_add( (wts.buffer.as_ref(), 0), (acc.buffer.as_ref(), 0), ], - &[mt_i, hid_i, uscale_f], + &[ffai_i, hid_i, uscale_f], grid, block, 0, @@ -4780,7 +4780,7 @@ pub fn relu2_scale_f16(dev: &dyn Device, x: &Tensor, scale: f32) -> Result Kernel { @@ -4813,9 +4813,9 @@ fn unary_act_kernel(name: &str, dtype: DType, kind: ActKind) -> Kernel { /// `acc[i] += x[i] · s[i]`. Lets the MoE expert sum stay ON-DEVICE — each /// expert's `down` output is folded into `acc` on the GPU (one final download /// per layer instead of one per expert). `s` is the per-expert weight broadcast -/// to `[len]`. Dispatches `mt_fma_inplace`. +/// to `[len]`. Dispatches `ffai_fma_inplace`. pub fn fma_inplace(dev: &dyn Device, acc: &Tensor, x: &Tensor, s: &Tensor) -> Result<()> { - let mut k = Kernel::new("mt_fma_inplace"); + let mut k = Kernel::new("ffai_fma_inplace"); for (pname, is_out) in [("acc", true), ("x", false), ("s", false)] { k.params.push(Param { name: pname.into(), dtype: acc.dtype, shape: Shape::scalar(), is_output: is_out, kind: ParamKind::Tensor }); } @@ -4839,7 +4839,7 @@ pub fn fma_inplace(dev: &dyn Device, acc: &Tensor, x: &Tensor, s: &Tensor) -> Re /// CLIP towers and the GELU-MLP LLM families. pub fn gelu(dev: &dyn Device, x: &Tensor) -> Result { let out = Tensor::empty(dev, x.shape.clone(), x.dtype)?; - let k = unary_act_kernel("mt_gelu", x.dtype, ActKind::Gelu); + let k = unary_act_kernel("ffai_gelu", x.dtype, ActKind::Gelu); let n = x.elem_count() as u32; let grid = Grid::d1(n.div_ceil(256), 256); dev.dispatch( @@ -4852,7 +4852,7 @@ pub fn gelu(dev: &dyn Device, x: &Tensor) -> Result { /// LayerNorm `out = (x - mean) / sqrt(var + eps) * w + b`, normalized over the /// last dim. Mean-subtracting + bias (unlike RMSNorm). Dispatches the -/// Reduction-mode `mt_layer_norm` (one threadgroup per row, block = n/4 — needs +/// Reduction-mode `ffai_layer_norm` (one threadgroup per row, block = n/4 — needs /// the row width `n` divisible by `4·lsize`; ViT/SigLIP widths are). Used by /// every transformer with LayerNorm (vision towers, BERT-style, GPT-2, …). pub fn layer_norm(dev: &dyn Device, x: &Tensor, weight: &Tensor, bias: &Tensor, eps: f32) -> Result { @@ -4860,7 +4860,7 @@ pub fn layer_norm(dev: &dyn Device, x: &Tensor, weight: &Tensor, bias: &Tensor, let rows = x.elem_count() / n; let out = Tensor::empty(dev, x.shape.clone(), x.dtype)?; let eps_buf = scalar_buf(dev, eps.to_bits())?; - let k = lookup("mt_layer_norm", x.dtype)?; + let k = lookup("ffai_layer_norm", x.dtype)?; let grid = Grid { grid: [rows as u32, 1, 1], block: [(n / 4) as u32, 1, 1] }; dev.dispatch( &k, @@ -4877,14 +4877,14 @@ pub fn layer_norm(dev: &dyn Device, x: &Tensor, weight: &Tensor, bias: &Tensor, Ok(out) } -/// Fused SwiGLU `out = silu(gate) * up`, elementwise. Dispatches `mt_swiglu`. +/// Fused SwiGLU `out = silu(gate) * up`, elementwise. Dispatches `ffai_swiglu`. pub fn swiglu(dev: &dyn Device, gate: &Tensor, up: &Tensor) -> Result { if gate.shape != up.shape { return Err(Error::Msg("swiglu: gate/up shape mismatch".into())); } elementwise_kernel( dev, - "mt_swiglu", + "ffai_swiglu", gate.dtype, gate.shape.clone(), vec![Binding::Buffer(gate.buffer.clone()), Binding::Buffer(up.buffer.clone())], @@ -4899,7 +4899,7 @@ pub fn gather(dev: &dyn Device, table: &Tensor, indices: &Tensor) -> Result Result { - let k = lookup("mt_partial_rope", qk.dtype)?; + let k = lookup("ffai_partial_rope", qk.dtype)?; let u = |x: u32| Binding::Scalar(x.to_le_bytes().to_vec()); let f = |x: f32| Binding::Scalar(x.to_le_bytes().to_vec()); // Bind qk as both input and output → in-place (nope dims preserved). @@ -4961,7 +4961,7 @@ pub fn dsv4_partial_rope( Ok(Tensor::new(qk.buffer.clone(), qk.shape.clone(), qk.dtype)) } -/// Softmax over the last dim, row-wise. Dispatches `mt_softmax`. Row width +/// Softmax over the last dim, row-wise. Dispatches `ffai_softmax`. Row width /// `n` must be a multiple of 1024 (the kernel's 4-elems/thread loop). pub fn softmax(dev: &dyn Device, x: &Tensor) -> Result { let n = *x.shape.last().ok_or_else(|| Error::Msg("softmax: scalar input".into()))?; @@ -4969,7 +4969,7 @@ pub fn softmax(dev: &dyn Device, x: &Tensor) -> Result { return Err(Error::Msg(format!("softmax: row width {n} must be a multiple of 1024"))); } let rows = x.elem_count() / n; - let k = lookup("mt_softmax", x.dtype)?; + let k = lookup("ffai_softmax", x.dtype)?; let out = Tensor::empty(dev, x.shape.clone(), x.dtype)?; let grid = Grid { grid: [rows as u32, 1, 1], block: [256, 1, 1] }; dev.dispatch( @@ -5002,9 +5002,9 @@ pub fn rope_llama( let head_dim = *qk.shape.last().ok_or_else(|| Error::Msg("rope: scalar input".into()))?; let n_heads = qk.elem_count() / head_dim; let half = head_dim / 2; - let k = lookup("mt_rope_banded", qk.dtype)?; + let k = lookup("ffai_rope_banded", qk.dtype)?; let out = Tensor::empty(dev, qk.shape.clone(), qk.dtype)?; - // mt_rope_banded reads `position` from a `positions: [T]` u32 buffer and + // ffai_rope_banded reads `position` from a `positions: [T]` u32 buffer and // selects the row via an outer grid axis; decode is the `T = 1` case (a // 1-element positions buffer). `row_stride` is the element stride between // consecutive rows (= n_heads*head_dim for a contiguous, non-fused-QKV qk). @@ -5058,7 +5058,7 @@ pub fn rope_llama_many( let t = qk.elem_count() / (n_heads * head_dim); let half = head_dim / 2; let k = cached_ir("ffai_rope_llama_many", qk.dtype, || { - let mut kk = ffai_kernels_std::kernels::rope::rope_banded::mt_rope_banded::kernel_ir_for(qk.dtype); + let mut kk = ffai_kernels_std::kernels::rope::rope_banded::ffai_rope_banded::kernel_ir_for(qk.dtype); kk.mode = KernelMode::Grid3D; kk }); @@ -5107,7 +5107,7 @@ pub fn kv_append_many( use ffai_kernels_core::ir::KernelMode; let t = src.elem_count() / (n_kv_heads * head_dim); let k = cached_ir("kv_cache_update_many", src.dtype, || { - let mut kk = ffai_kernels_std::kernels::kv_cache::update_many::mt_kv_cache_update_many::kernel_ir_for(src.dtype); + let mut kk = ffai_kernels_std::kernels::kv_cache::update_many::ffai_kv_cache_update_many::kernel_ir_for(src.dtype); kk.mode = KernelMode::Grid3D; kk }); @@ -5269,7 +5269,7 @@ pub fn mamba_split_proj( ) -> Result<(Tensor, Tensor, Tensor)> { let kernel = cached_ir("mamba_split_proj", DType::F32, || { use ffai_kernels_core::ir::KernelMode; - let mut k = ffai_kernels_std::kernels::ssm::scan::mt_mamba_split_proj::kernel_ir_for(); + let mut k = ffai_kernels_std::kernels::ssm::scan::ffai_mamba_split_proj::kernel_ir_for(); k.mode = KernelMode::Grid3D; k }); @@ -5317,7 +5317,7 @@ pub fn mamba_split_conv( ) -> Result<(Tensor, Tensor, Tensor)> { let kernel = cached_ir("mamba_split_conv", DType::F32, || { use ffai_kernels_core::ir::KernelMode; - let mut k = ffai_kernels_std::kernels::ssm::scan::mt_mamba_split_conv::kernel_ir_for(); + let mut k = ffai_kernels_std::kernels::ssm::scan::ffai_mamba_split_conv::kernel_ir_for(); k.mode = KernelMode::Grid3D; k }); @@ -5361,7 +5361,7 @@ pub fn gated_group_rmsnorm_batched( let ng = di / gs; let kernel = cached_ir("gated_group_rmsnorm_batched", DType::F32, || { use ffai_kernels_core::ir::KernelMode; - let mut k = ffai_kernels_std::kernels::ssm::scan::mt_gated_group_rmsnorm_batched::kernel_ir_for(); + let mut k = ffai_kernels_std::kernels::ssm::scan::ffai_gated_group_rmsnorm_batched::kernel_ir_for(); k.mode = KernelMode::Reduction; k }); @@ -6674,7 +6674,7 @@ extern "C" __global__ void moe_fp4_actq_pack( const int* __restrict__ gend, const float* __restrict__ gA, unsigned* __restrict__ Ap, unsigned* __restrict__ Asc, const unsigned* __restrict__ st, - int maxt, int K, int mt_total, int act_mode, float act_scale, int use_st) + int maxt, int K, int ffai_total, int act_mode, float act_scale, int use_st) { // Warp-per-(tile,kt): stage the 16x64 f16 tile to smem with coalesced // half2 loads (each lane: 2 rows x 16B), compute the 64 block-16 amaxes @@ -6695,7 +6695,7 @@ extern "C" __global__ void moe_fp4_actq_pack( { int r0=lane>>1, c0=(lane&1)*32; // each lane: row r0, cols c0..c0+31 int row=base+r0; - bool live=(r0 < end-base)&&(rowshared copy (per-warp staging; no block syncs). -__device__ __forceinline__ void mt_cp16(void* sm, const void* gm){ +__device__ __forceinline__ void ffai_cp16(void* sm, const void* gm){ unsigned saddr=(unsigned)__cvta_generic_to_shared(sm); asm volatile("cp.async.cg.shared.global [%0], [%1], 16;"::"r"(saddr),"l"(gm)); } @@ -6772,7 +6772,7 @@ extern "C" __global__ void moe_fp4_grouped_mma_x4( const int* __restrict__ gend, const float* __restrict__ gA, const float* __restrict__ gw, __half* __restrict__ out, unsigned* amax_out, - int N, int K, int mt_total, int am_mode, float am_scale) + int N, int K, int ffai_total, int am_mode, float am_scale) { // Per-warp cp.async 2-stage smem W pipeline + A register prefetch: the // direct global->register W loads left the mma latency-bound (~42 TF); @@ -6799,12 +6799,12 @@ extern "C" __global__ void moe_fp4_grouped_mma_x4( int c0=lane, c1=lane+32; \ int j0=c0>>4, j1=c1>>4; \ int nt_a=nt0+j0, nt_b=nt0+j1; \ - if(nt_a>1, hf=lane&1; int nt_s=nt0+js; \ - if(nt_s Result<()> { let kern = cached_ir("ffai_gemm_batched", DType::F32, || { - let mut k = ffai_kernels_std::kernels::gemm::dense::mt_gemm_batched::kernel_ir_for(DType::F32); + let mut k = ffai_kernels_std::kernels::gemm::dense::ffai_gemm_batched::kernel_ir_for(DType::F32); k.mode = KernelMode::Reduction; k }); @@ -317,15 +317,15 @@ pub fn ssm_prefill_scan_ssd_portable( fn build_ssd_ir(name: &str) -> ffai_core::Kernel { use ffai_kernels_std::kernels::ssm; match name { - "ssd_lcs" => ssm::scan::mt_ssd_lcs::kernel_ir_for(), - "ssd_gather_bc" => ssm::scan::mt_ssd_gather_bc::kernel_ir_for(), - "ssd_xt" => ssm::scan::mt_ssd_xt::kernel_ir_for(), - "ssd_mmask" => ssm::scan::mt_ssd_mmask::kernel_ir_for(), - "ssd_bdt" => ssm::scan::mt_ssd_bdt::kernel_ir_for(), - "ssd_recur" => ssm::scan::mt_ssd_recur::kernel_ir_for(), - "ssd_combine" => ssm::scan::mt_ssd_combine::kernel_ir_for(), - "ssd_g1_cb" => ssm::scan::mt_ssd_g1_cb::kernel_ir_for(), - "ssd_g4_cs" => ssm::scan::mt_ssd_g4_cs::kernel_ir_for(), + "ssd_lcs" => ssm::scan::ffai_ssd_lcs::kernel_ir_for(), + "ssd_gather_bc" => ssm::scan::ffai_ssd_gather_bc::kernel_ir_for(), + "ssd_xt" => ssm::scan::ffai_ssd_xt::kernel_ir_for(), + "ssd_mmask" => ssm::scan::ffai_ssd_mmask::kernel_ir_for(), + "ssd_bdt" => ssm::scan::ffai_ssd_bdt::kernel_ir_for(), + "ssd_recur" => ssm::scan::ffai_ssd_recur::kernel_ir_for(), + "ssd_combine" => ssm::scan::ffai_ssd_combine::kernel_ir_for(), + "ssd_g1_cb" => ssm::scan::ffai_ssd_g1_cb::kernel_ir_for(), + "ssd_g4_cs" => ssm::scan::ffai_ssd_g4_cs::kernel_ir_for(), other => panic!("build_ssd_ir: unknown ssd kernel {other}"), } }