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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 6 additions & 6 deletions rust/Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

4 changes: 2 additions & 2 deletions rust/crates/backends/ffai-cuda/src/imp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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())
}

Expand Down
14 changes: 7 additions & 7 deletions rust/crates/backends/ffai-cuda/tests/cuda_smoke.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -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;
Expand All @@ -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 {
Expand Down Expand Up @@ -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<f32> = (0..1024).map(|i| (i % 13) as f32 * 0.05).collect();
Expand All @@ -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;
Expand Down Expand Up @@ -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
Expand Down
6 changes: 3 additions & 3 deletions rust/crates/backends/ffai-cuda/tests/ssm_test.rs
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -13,7 +13,7 @@ fn fill(n:usize,s:usize)->Vec<f32>{(0..n).map(|i|(((i*7+s*131)%89) as f32-44.0)*
fn tn(d:&dyn Device,v:&[f32],sh:Vec<usize>)->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;
Expand Down Expand Up @@ -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.");
}
4 changes: 2 additions & 2 deletions rust/crates/backends/ffai-metal/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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())
}

Expand Down
32 changes: 31 additions & 1 deletion rust/crates/backends/ffai-metal/tests/gemv_q8.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,34 @@ use ffai_metal::MetalDevice;
fn tb_f32(v: &[f32]) -> Vec<u8> { v.iter().flat_map(|x| x.to_le_bytes()).collect() }
fn tb_u32(v: &[u32]) -> Vec<u8> { v.iter().flat_map(|x| x.to_le_bytes()).collect() }
fn fb(b: &[u8]) -> Vec<f32> { 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: 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();
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<u8> { v.iter().flat_map(|&f| f32_to_f16(f).to_le_bytes()).collect() }

#[test]
fn gemv_q8_matches_cpu_dequant_dot() {
Expand Down Expand Up @@ -98,6 +126,8 @@ fn gemv_q4_matches_cpu_dequant() {
let w: Vec<f32> = (0..m * k).map(|_| rng()).collect();
let x: Vec<f32> = (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<f32> = 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;
Expand All @@ -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();
Expand Down
6 changes: 3 additions & 3 deletions rust/crates/backends/ffai-metal/tests/ssm_test.rs
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -12,7 +12,7 @@ fn fill(n:usize,s:usize)->Vec<f32>{(0..n).map(|i|(((i*7+s*131)%89) as f32-44.0)*
fn tn(d:&dyn Device,v:&[f32],sh:Vec<usize>)->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;
Expand Down Expand Up @@ -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.");
}
Expand Down
4 changes: 2 additions & 2 deletions rust/crates/backends/ffai-vulkan/src/imp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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())
}

Expand Down
12 changes: 6 additions & 6 deletions rust/crates/ffai-modeltests/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
Loading