#[cfg(not(memra_cutlass))]
fn main() {
eprintln!("build with MEMRA_CUTLASS=1 on sm_120a");
}
#[cfg(memra_cutlass)]
fn main() -> Result<(), Box<dyn std::error::Error>> {
use memra_engine::Engine;
let eng = Engine::new(0)?;
let shapes: &[(usize, usize, usize, &str)] = &[
(1, 20480, 4096, "gate+up all-8-experts m=1"),
(1, 4096, 5120, "down all-8-experts m=1 (k=8*640)"),
(2, 20480, 4096, "gate+up m=2 (verify T=2)"),
(8, 20480, 4096, "gate+up m=8 (batch B=8)"),
(1, 6144, 4096, "qkv-class m=1 (bf16 comparison shape)"),
(512, 6144, 4096, "PREFILL qkv chunk m=512"),
(512, 4096, 4096, "PREFILL o_proj chunk m=512"),
(16, 2560, 4096, "PREFILL expert gate+up tile m=16"),
(16, 4096, 1280, "PREFILL expert down tile m=16"),
(512, 2560, 4096, "PREFILL expert gate+up dense-equiv m=512"),
];
for &(m, n, k, label) in shapes {
let us = time_shape(&eng, m, n, k)?;
let bytes = (n * k) as f64 / 2.0 + (n * k) as f64 / 16.0;
let tflops = 2.0 * m as f64 * n as f64 * k as f64 / us / 1e6;
println!(
"[{label}] m={m} n={n} k={k}: {us:.1} us/call weightBW={:.2} TB/s {tflops:.1} TFLOP/s",
bytes / us / 1e6
);
}
Ok(())
}
#[cfg(memra_cutlass)]
fn time_shape(
eng: &memra_engine::Engine,
m: usize,
n: usize,
k: usize,
) -> Result<f64, Box<dyn std::error::Error>> {
let a_h = vec![0.25f32; m * k];
let b_h = vec![0.5f32; n * k];
let a_d = eng.htod(&a_h)?;
let b_d = eng.htod(&b_h)?;
let mut a_packed = eng.alloc_u8(m * k / 2)?;
let mut a_sf_lin = eng.alloc_u8(m * k / 16)?;
let mut b_packed = eng.alloc_u8(n * k / 2)?;
let mut b_sf_lin = eng.alloc_u8(n * k / 16)?;
eng.cutlass_nvfp4_quant_ref(&a_d, &mut a_packed, &mut a_sf_lin, m, k)?;
eng.cutlass_nvfp4_quant_ref(&b_d, &mut b_packed, &mut b_sf_lin, n, k)?;
let sfa_bytes = eng.cutlass_sfa_size(m, k);
let sfb_bytes = eng.cutlass_sfb_size(n, k);
let mut a_sf_sw = eng.alloc_u8(sfa_bytes)?;
let mut b_sf_sw = eng.alloc_u8(sfb_bytes)?;
eng.cutlass_repack_sfa(&a_sf_lin, &mut a_sf_sw, m, k)?;
eng.cutlass_repack_sfb(&b_sf_lin, &mut b_sf_sw, n, k)?;
let alpha_d = eng.htod(&[1.0f32])?;
let ws_bytes = eng.cutlass_fp4_workspace_size(m, n, k);
let mut workspace = eng.alloc_u8(ws_bytes.max(1))?;
let mut d_d = eng.htod(&vec![0f32; m * n])?;
for _ in 0..20 {
eng.cutlass_fp4_gemm_raw(
&a_packed,
&b_packed,
&a_sf_sw,
&b_sf_sw,
&alpha_d,
&mut d_d,
m,
n,
k,
&mut workspace,
)?;
}
eng.stream().synchronize()?;
let reps = 200;
let t0 = std::time::Instant::now();
for _ in 0..reps {
eng.cutlass_fp4_gemm_raw(
&a_packed,
&b_packed,
&a_sf_sw,
&b_sf_sw,
&alpha_d,
&mut d_d,
m,
n,
k,
&mut workspace,
)?;
}
eng.stream().synchronize()?;
Ok(t0.elapsed().as_secs_f64() * 1e6 / reps as f64)
}