use hanzo_kernel::attn::{sdpa_blk, sdpa_blk_run, sdpa_ref};
use hanzo_kernel::norm::rms_norm_run;
use hanzo_kernel::prelude::*;
use hanzo_kernel::quant::{
gemv, gemv_run, gen_moe_q4k, gen_moe_q6k, gen_q4k, gen_q8_0_packed, matvec_q4k_bench,
matvec_q4k_dp4a_blk, matvec_q4k_dp4a_blk_run, matvec_q4k_f32_blk_run, matvec_q4k_ref,
matvec_q4k_run, matvec_q8_0_packed_ref, matvec_q8_0_packed_run, matvec_q8_0_packed_sg_run,
matvec_q8_bench, matvec_q8_dp4a_blk_run, matvec_q8_dp4a_i8_run, matvec_q8_dp4a_ref,
matvec_q8_ref, matvec_q8_run, moe_matvec_q4k_bench, moe_matvec_q4k_blk_bench,
moe_matvec_q4k_blk_run, moe_matvec_q4k_dp4a_blk, moe_matvec_q4k_dp4a_blk_run,
moe_matvec_q4k_ref, moe_matvec_q4k_run, moe_matvec_q6k_bench, moe_matvec_q6k_blk_bench,
moe_matvec_q6k_blk_run, moe_matvec_q6k_dp4a_blk, moe_matvec_q6k_dp4a_blk_run,
moe_matvec_q6k_ref, moe_matvec_q6k_run, moe_route, moe_route_ref, moe_route_run, pack_q4k,
quant_act_q8_cpu, QK8_0,
};
use std::time::Instant;
fn check_dp4a<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
rows: usize,
k: usize,
coop: bool,
) {
let mut s = 0x9E3779B9_7F4A7C15u64;
let mut nxt = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let wq8: Vec<i8> = (0..rows * k).map(|_| (nxt() % 255) as i8).collect();
let xq8: Vec<i8> = (0..k).map(|_| (nxt() % 255) as i8).collect();
let wd: Vec<f32> = (0..rows * k / 32)
.map(|_| (nxt() % 1000) as f32 / 8000.0 + 0.01)
.collect();
let wq32: Vec<i32> = wq8.iter().map(|&x| x as i32).collect();
let xq32: Vec<i32> = xq8.iter().map(|&x| x as i32).collect();
let reference = matvec_q8_dp4a_ref(&wq32, &xq32, &wd, rows, k);
let real_bytes = (rows * k) as f64; let flop = 2.0 * rows as f64 * k as f64;
let mut report = |tag: &str, (out, ms): (Vec<f32>, f64)| {
let rel = max_rel(&reference, &out);
println!(
"[{:<7}] dp4a/{:<6} {}x{} max_rel={:.2e} {} {:.3} ms {:.0} GB/s {:.0} GFLOP/s",
name,
tag,
rows,
k,
rel,
if rel < 2e-2 {
"MATCH ✓"
} else {
"MISMATCH ✗"
},
ms,
real_bytes / (ms * 1e6),
flop / (ms * 1e6)
);
};
report(
"i8pack",
matvec_q8_dp4a_i8_run(client, &wq8, &xq8, &wd, rows, k, 50),
);
if coop {
for nt in [64usize, 128, 256] {
report(
&format!("blk{nt}"),
matvec_q8_dp4a_blk_run(client, &wq8, &xq8, &wd, rows, k, nt, 50),
);
}
}
}
fn maxrel(a: &[f32], b: &[f32]) -> f32 {
let mut m = 0f32;
for (x, y) in a.iter().zip(b.iter()) {
m = m.max((x - y).abs() / x.abs().max(1e-6));
}
m
}
fn scalerel(a: &[f32], b: &[f32]) -> f32 {
let scale = a.iter().fold(0f32, |m, x| m.max(x.abs())).max(1e-20);
a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs())) / scale
}
fn check_q4k<R: Runtime>(name: &str, client: &ComputeClient<R>, rows: usize, k: usize) {
let (wqs, wsc, wd, wdm, x) = gen_q4k(rows, k);
let reference = matvec_q4k_ref(&wqs, &wsc, &wd, &wdm, &x, rows, k);
let got = matvec_q4k_run::<R>(client, &wqs, &wsc, &wd, &wdm, &x, rows, k);
let rel = maxrel(&reference, &got);
let ok = rel < 3e-3;
let ms = matvec_q4k_bench::<R>(client, &wqs, &wsc, &wd, &wdm, &x, rows, k, 50);
let wbytes = rows * (k / 256) * 144;
let gbps = wbytes as f64 / (ms * 1e6);
let gflops = 2.0 * rows as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] Q4_K {}x{} max_rel={:.2e} {} {:.3} ms {:.0} GB/s {:.0} GFLOP/s",
name,
rows,
k,
rel,
if ok { "BIT-EXACT ✓" } else { "MISMATCH ✗" },
ms,
gbps,
gflops
);
}
fn check_moe_q4k<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
e: usize,
n: usize,
slots: usize,
k: usize,
) {
let (wqs, wsc, wd, wdm, x, ids) = gen_moe_q4k(e, n, slots, k);
let reference = moe_matvec_q4k_ref(&wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k);
let got = moe_matvec_q4k_run::<R>(client, &wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k);
let rel = maxrel(&reference, &got);
let ok = rel < 3e-3;
let ms = moe_matvec_q4k_bench::<R>(client, &wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k, 50);
let wbytes = slots * n * (k / 256) * 144;
let gbps = wbytes as f64 / (ms * 1e6);
let gflops = 2.0 * (slots * n) as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] MoEQ4K E{} {}x{}x{} max_rel={:.2e} {} {:.3} ms {:.0} GB/s {:.0} GFLOP/s",
name,
e,
slots,
n,
k,
rel,
if ok { "BIT-EXACT ✓" } else { "MISMATCH ✗" },
ms,
gbps,
gflops
);
}
fn check_matvec_q4k_dp4a_blk<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
nout: usize,
k: usize,
nt: usize,
) {
let (wqs, wsc, wd, wdm, x) = gen_q4k(nout, k);
let reference = matvec_q4k_ref(&wqs, &wsc, &wd, &wdm, &x, nout, k);
let got = matvec_q4k_dp4a_blk_run::<R>(client, &wqs, &wsc, &wd, &wdm, &x, nout, k, nt);
let srel = scalerel(&reference, &got);
let ok = srel < 1e-2;
let packed = pack_q4k(&wqs, &wsc, &wd, &wdm);
let (xq, xs, xsum) = quant_act_q8_cpu(&x, 1, k);
let wh = client.create_from_slice(u32::as_bytes(&packed));
let xqh = client.create_from_slice(u32::as_bytes(&xq));
let xsh = client.create_from_slice(f32::as_bytes(&xs));
let xsumh = client.create_from_slice(f32::as_bytes(&xsum));
let meta = client.create_from_slice(u32::as_bytes(&[k as u32]));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; nout]));
let launch = |c: &ComputeClient<R>| unsafe {
matvec_q4k_dp4a_blk::launch_unchecked::<f32, R>(
c,
Grid::Static(nout as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(wh.clone(), packed.len()),
ArrayArg::from_raw_parts(xqh.clone(), xq.len()),
ArrayArg::from_raw_parts(xsh.clone(), xs.len()),
ArrayArg::from_raw_parts(xsumh.clone(), xsum.len()),
ArrayArg::from_raw_parts(oh.clone(), nout),
ArrayArg::from_raw_parts(meta.clone(), 1),
nt,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t0 = Instant::now();
for _ in 0..100 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let ms = t0.elapsed().as_secs_f64() * 1e3 / 100.0;
let gbps = (nout * (k / 256) * 144) as f64 / (ms * 1e6);
println!(
"[{:<7}] Q4Kdp4a {}x{} nt={:<3} scale_rel={:.2e} {} {:.4} ms {:.0} GB/s",
name,
nout,
k,
nt,
srel,
if ok { "MATCH ✓" } else { "MISMATCH ✗" },
ms,
gbps
);
}
fn check_moe_q4k_blk<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
e: usize,
n: usize,
slots: usize,
k: usize,
nt: usize,
) {
let (wqs, wsc, wd, wdm, x, ids) = gen_moe_q4k(e, n, slots, k);
let reference = moe_matvec_q4k_ref(&wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k);
let got = moe_matvec_q4k_blk_run::<R>(client, &wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k, nt);
let srel = scalerel(&reference, &got);
let ok = srel < 1e-3;
let ms =
moe_matvec_q4k_blk_bench::<R>(client, &wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k, nt, 50);
let wbytes = slots * n * (k / 256) * 144;
let gbps = wbytes as f64 / (ms * 1e6);
let gflops = 2.0 * (slots * n) as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] MoEQ4Kb E{} {}x{}x{} nt={:<3} scale_rel={:.2e} {} {:.3} ms {:.0} GB/s {:.0} GFLOP/s",
name, e, slots, n, k, nt, srel,
if ok { "MATCH ✓" } else { "MISMATCH ✗" }, ms, gbps, gflops
);
}
fn check_moe_q4k_dp4a_blk<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
e: usize,
n: usize,
slots: usize,
k: usize,
nt: usize,
) {
let (wqs, wsc, wd, wdm, x, ids) = gen_moe_q4k(e, n, slots, k);
let reference = moe_matvec_q4k_ref(&wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k);
let got =
moe_matvec_q4k_dp4a_blk_run::<R>(client, &wqs, &wsc, &wd, &wdm, &x, &ids, slots, n, k, nt);
let srel = scalerel(&reference, &got);
let ok = srel < 1e-2;
let (xq, xs, xsum) = quant_act_q8_cpu(&x, slots, k);
let qh = client.create_from_slice(u32::as_bytes(&wqs));
let sh = client.create_from_slice(u32::as_bytes(&wsc));
let dh = client.create_from_slice(f32::as_bytes(&wd));
let mh = client.create_from_slice(f32::as_bytes(&wdm));
let xqh = client.create_from_slice(u32::as_bytes(&xq));
let xsh = client.create_from_slice(f32::as_bytes(&xs));
let xsumh = client.create_from_slice(f32::as_bytes(&xsum));
let ih = client.create_from_slice(u32::as_bytes(&ids));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n]));
let launch = |c: &ComputeClient<R>| unsafe {
moe_matvec_q4k_dp4a_blk::launch_unchecked::<f32, R>(
c,
Grid::Static((slots * n) as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(qh.clone(), wqs.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(mh.clone(), wdm.len()),
ArrayArg::from_raw_parts(xqh.clone(), xq.len()),
ArrayArg::from_raw_parts(xsh.clone(), xs.len()),
ArrayArg::from_raw_parts(xsumh.clone(), xsum.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
nt,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t0 = Instant::now();
for _ in 0..50 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let ms = t0.elapsed().as_secs_f64() * 1e3 / 50.0;
let gbps = (slots * n * (k / 256) * 144) as f64 / (ms * 1e6);
let gflops = 2.0 * (slots * n) as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] MoEdp4a E{} {}x{}x{} nt={:<3} scale_rel={:.2e} {} {:.3} ms {:.0} GB/s {:.0} GFLOP/s",
name, e, slots, n, k, nt, srel,
if ok { "MATCH ✓" } else { "MISMATCH ✗" }, ms, gbps, gflops
);
}
fn check_moe_q6k<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
e: usize,
n: usize,
slots: usize,
k: usize,
) {
let (wql, wqh, wsc, wd, x, ids) = gen_moe_q6k(e, n, slots, k);
let reference = moe_matvec_q6k_ref(&wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k);
let got = moe_matvec_q6k_run::<R>(client, &wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k);
let srel = scalerel(&reference, &got);
let ok = srel < 1e-3;
let ms = moe_matvec_q6k_bench::<R>(client, &wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k, 50);
let wbytes = slots * n * (k / 256) * 210; let gbps = wbytes as f64 / (ms * 1e6);
let gflops = 2.0 * (slots * n) as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] MoEQ6K E{} {}x{}x{} scale_rel={:.2e} {} {:.3} ms {:.0} GB/s {:.0} GFLOP/s",
name,
e,
slots,
n,
k,
srel,
if ok { "MATCH ✓" } else { "MISMATCH ✗" },
ms,
gbps,
gflops
);
}
fn check_moe_q6k_dp4a_blk<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
e: usize,
n: usize,
slots: usize,
k: usize,
nt: usize,
) {
let (wql, wqh, wsc, wd, x, ids) = gen_moe_q6k(e, n, slots, k);
let reference = moe_matvec_q6k_ref(&wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k);
let got =
moe_matvec_q6k_dp4a_blk_run::<R>(client, &wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k, nt);
let srel = scalerel(&reference, &got);
let ok = srel < 1e-2;
let (xq, xs, _) = quant_act_q8_cpu(&x, slots, k);
let qlh = client.create_from_slice(u32::as_bytes(&wql));
let qhh = client.create_from_slice(u32::as_bytes(&wqh));
let sh = client.create_from_slice(u32::as_bytes(&wsc));
let dh = client.create_from_slice(f32::as_bytes(&wd));
let xqh = client.create_from_slice(u32::as_bytes(&xq));
let xsh = client.create_from_slice(f32::as_bytes(&xs));
let ih = client.create_from_slice(u32::as_bytes(&ids));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; slots * n]));
let launch = |c: &ComputeClient<R>| unsafe {
moe_matvec_q6k_dp4a_blk::launch_unchecked::<f32, R>(
c,
Grid::Static((slots * n) as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(qlh.clone(), wql.len()),
ArrayArg::from_raw_parts(qhh.clone(), wqh.len()),
ArrayArg::from_raw_parts(sh.clone(), wsc.len()),
ArrayArg::from_raw_parts(dh.clone(), wd.len()),
ArrayArg::from_raw_parts(xqh.clone(), xq.len()),
ArrayArg::from_raw_parts(xsh.clone(), xs.len()),
ArrayArg::from_raw_parts(ih.clone(), ids.len()),
ArrayArg::from_raw_parts(oh.clone(), slots * n),
n,
k,
nt,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t0 = Instant::now();
for _ in 0..50 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let ms = t0.elapsed().as_secs_f64() * 1e3 / 50.0;
let gbps = (slots * n * (k / 256) * 210) as f64 / (ms * 1e6);
let gflops = 2.0 * (slots * n) as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] MoEq6dp4a E{} {}x{}x{} nt={:<3} scale_rel={:.2e} {} {:.3} ms {:.0} GB/s {:.0} GFLOP/s",
name, e, slots, n, k, nt, srel,
if ok { "MATCH ✓" } else { "MISMATCH ✗" }, ms, gbps, gflops
);
}
fn check_moe_q6k_blk<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
e: usize,
n: usize,
slots: usize,
k: usize,
nt: usize,
) {
let (wql, wqh, wsc, wd, x, ids) = gen_moe_q6k(e, n, slots, k);
let reference = moe_matvec_q6k_ref(&wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k);
let got = moe_matvec_q6k_blk_run::<R>(client, &wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k, nt);
let srel = scalerel(&reference, &got);
let ok = srel < 1e-3;
let ms =
moe_matvec_q6k_blk_bench::<R>(client, &wql, &wqh, &wsc, &wd, &x, &ids, slots, n, k, nt, 50);
let wbytes = slots * n * (k / 256) * 210;
let gbps = wbytes as f64 / (ms * 1e6);
let gflops = 2.0 * (slots * n) as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] MoEQ6Kb E{} {}x{}x{} nt={:<3} scale_rel={:.2e} {} {:.3} ms {:.0} GB/s {:.0} GFLOP/s",
name, e, slots, n, k, nt, srel,
if ok { "MATCH ✓" } else { "MISMATCH ✗" }, ms, gbps, gflops
);
}
fn check_moe_route<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
ntok: usize,
n_experts: usize,
topk: usize,
nt: usize,
) {
let mut s = 0x9E3779B97F4A7C15u64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let logits: Vec<f32> = (0..ntok * n_experts)
.map(|_| (next() % 4000) as f32 / 1000.0 - 2.0)
.collect();
let (ids_ref, w_ref) = moe_route_ref(&logits, ntok, n_experts, topk);
let (ids_got, w_got) = moe_route_run::<R>(client, &logits, ntok, n_experts, topk, nt);
let mut ties = 0usize;
for tok in 0..ntok {
let mut a: Vec<u32> = ids_ref[tok * topk..(tok + 1) * topk].to_vec();
a.sort_unstable();
let mut b: Vec<u32> = ids_got[tok * topk..(tok + 1) * topk].to_vec();
b.sort_unstable();
if a != b {
ties += 1;
}
}
let wrel = scalerel(&w_ref, &w_got);
let ok = wrel < 1e-4;
let lh = client.create_from_slice(f32::as_bytes(&logits));
let ih = client.create_from_slice(u32::as_bytes(&vec![0u32; ntok * topk]));
let wh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; ntok * topk]));
let launch = |c: &ComputeClient<R>| unsafe {
moe_route::launch_unchecked::<f32, R>(
c,
Grid::Static(ntok as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(lh.clone(), logits.len()),
ArrayArg::from_raw_parts(ih.clone(), ntok * topk),
ArrayArg::from_raw_parts(wh.clone(), ntok * topk),
n_experts,
topk,
nt,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(ih.clone());
let t0 = std::time::Instant::now();
for _ in 0..200 {
launch(client);
}
let _ = client.read_one_unchecked(ih.clone());
let ms = t0.elapsed().as_secs_f64() * 1e3 / 200.0;
println!(
"[{:<7}] MoERoute ntok={} E{} topk={} nt={:<3} w_rel={:.2e} ties={}/{} {} {:.4} ms/batch",
name,
ntok,
n_experts,
topk,
nt,
wrel,
ties,
ntok,
if ok { "MATCH ✓" } else { "MISMATCH ✗" },
ms
);
}
fn check_sdpa_blk<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
n_heads: usize,
n_kv: usize,
seq_q: usize,
seq_k: usize,
kv_seq_pad: usize,
d: usize,
causal: bool,
nt: usize,
) {
let mut s = 0x2545F4914F6CDD1Du64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let mut rnd = |_| (next() % 2000) as f32 / 1000.0 - 1.0;
let q: Vec<f32> = (0..n_heads * seq_q * d).map(&mut rnd).collect();
let k: Vec<f32> = (0..n_kv * seq_k * d).map(&mut rnd).collect();
let v: Vec<f32> = (0..n_kv * seq_k * d).map(&mut rnd).collect();
let want = sdpa_ref(&q, &k, &v, n_heads, n_kv, seq_q, seq_k, d, causal);
let pad = |src: &[f32]| -> Vec<f32> {
let mut buf = vec![0.0f32; n_kv * kv_seq_pad * d];
for kvh in 0..n_kv {
let (s0, d0) = (kvh * seq_k * d, kvh * kv_seq_pad * d);
buf[d0..d0 + seq_k * d].copy_from_slice(&src[s0..s0 + seq_k * d]);
}
buf
};
let kp = pad(&k);
let vp = pad(&v);
let got = sdpa_blk_run::<R>(
client, &q, &kp, &vp, n_heads, n_kv, seq_q, seq_k, kv_seq_pad, d, causal, nt,
);
let srel = scalerel(&want, &got);
let mrel = max_rel(&want, &got);
let ok = srel < 1e-4;
let scale = 1.0f32 / (d as f32).sqrt();
let meta = [
seq_q as u32,
seq_k as u32,
n_heads as u32,
n_kv as u32,
causal as u32,
(n_kv * kv_seq_pad * d) as u32,
(kv_seq_pad * d) as u32,
d as u32,
];
let qh = client.create_from_slice(f32::as_bytes(&q));
let kh = client.create_from_slice(f32::as_bytes(&kp));
let vh = client.create_from_slice(f32::as_bytes(&vp));
let sh = client.create_from_slice(f32::as_bytes(&[scale]));
let mh = client.create_from_slice(u32::as_bytes(&meta));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; n_heads * seq_q * d]));
let launch = |c: &ComputeClient<R>| unsafe {
sdpa_blk::launch_unchecked::<f32, R>(
c,
Grid::Static((n_heads * seq_q) as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(qh.clone(), q.len()),
ArrayArg::from_raw_parts(kh.clone(), kp.len()),
ArrayArg::from_raw_parts(vh.clone(), vp.len()),
ArrayArg::from_raw_parts(oh.clone(), n_heads * seq_q * d),
ArrayArg::from_raw_parts(sh.clone(), 1),
ArrayArg::from_raw_parts(mh.clone(), 8),
d,
nt,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t0 = Instant::now();
for _ in 0..200 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let ms = t0.elapsed().as_secs_f64() * 1e3 / 200.0;
let kv_bytes = (n_heads * seq_k * 2 * d * 4) as f64; let gbs = kv_bytes / (ms * 1e-3) / 1e9;
println!(
"[{:<7}] SDPAblk h{}kv{} q{} k{} d{} nt={:<3} scale_rel={:.2e} (max_rel={:.1e}) {} {:.4} ms {:.0} GB/s",
name, n_heads, n_kv, seq_q, seq_k, d, nt, srel, mrel,
if ok { "MATCH ✓" } else { "MISMATCH ✗" }, ms, gbs
);
}
fn check_gemv<R: Runtime>(name: &str, client: &ComputeClient<R>, n: usize, k: usize, nt: usize) {
let mut s = 0x9E3779B97F4A7C15u64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let w: Vec<f32> = (0..n * k)
.map(|_| (next() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let x: Vec<f32> = (0..k)
.map(|_| (next() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let want: Vec<f32> = (0..n)
.map(|r| (0..k).map(|i| w[r * k + i] * x[i]).sum())
.collect();
let got = gemv_run::<R>(client, &w, &x, n, k, nt);
let rel = scalerel(&want, &got);
let ok = rel < 1e-4;
let wh = client.create_from_slice(f32::as_bytes(&w));
let xh = client.create_from_slice(f32::as_bytes(&x));
let oh = client.create_from_slice(f32::as_bytes(&vec![0.0f32; n]));
let mh = client.create_from_slice(u32::as_bytes(&[k as u32]));
let launch = |c: &ComputeClient<R>| unsafe {
gemv::launch_unchecked::<f32, R>(
c,
Grid::Static(n as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(wh.clone(), w.len()),
ArrayArg::from_raw_parts(xh.clone(), x.len()),
ArrayArg::from_raw_parts(oh.clone(), n),
ArrayArg::from_raw_parts(mh.clone(), 1),
nt,
);
};
for _ in 0..3 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let t0 = Instant::now();
for _ in 0..200 {
launch(client);
}
let _ = client.read_one_unchecked(oh.clone());
let ms = t0.elapsed().as_secs_f64() * 1e3 / 200.0;
let gbs = (n * k * 4) as f64 / (ms * 1e-3) / 1e9;
println!(
"[{:<7}] GEMV {}x{} nt={:<3} scale_rel={:.2e} {} {:.4} ms {:.0} GB/s",
name,
n,
k,
nt,
rel,
if ok { "MATCH ✓" } else { "MISMATCH ✗" },
ms,
gbs
);
}
fn check_mmq_q4k<R: Runtime>(name: &str, client: &ComputeClient<R>, m: usize, n: usize, k: usize) {
use hanzo_kernel::mmq::{gen_mmq_q4k, mmq_q4k_ref, mmq_q4k_wmma_blk_run};
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(m, n, k);
let want = mmq_q4k_ref(&xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k);
let (got, ms) =
mmq_q4k_wmma_blk_run::<R>(client, &xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k, 50);
let rel = maxabs_over_max(&got, &want);
let gflops = 2.0 * m as f64 * n as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] MMQ-Q4K {}x{}x{} rel={:.2e} {} {:.3} ms {:.0} GFLOP/s",
name,
m,
n,
k,
rel,
if rel < 1e-2 {
"COOPMAT ✓"
} else {
"MISMATCH ✗"
},
ms,
gflops
);
}
fn check_mmq_q4k_rt<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
m: usize,
n: usize,
k: usize,
) {
use hanzo_kernel::mmq::{gen_mmq_q4k, mmq_q4k_ref, mmq_q4k_wmma_rt_run};
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(m, n, k);
let want = mmq_q4k_ref(&xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k);
let (got, ms) =
mmq_q4k_wmma_rt_run::<R>(client, &xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, m, n, k, 50);
let rel = maxabs_over_max(&got, &want);
let gflops = 2.0 * m as f64 * n as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] MMQ-Q4K-RT {}x{}x{} rel={:.2e} {} {:.3} ms {:.0} GFLOP/s",
name,
m,
n,
k,
rel,
if rel < 1e-2 {
"COOPMAT ✓"
} else {
"MISMATCH ✗"
},
ms,
gflops
);
}
fn check_mmq_q4k_id<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
t: usize,
topk: usize,
n_experts: usize,
n: usize,
k: usize,
) {
use hanzo_kernel::mmq::{gen_mmq_q4k, mmq_q4k_id_ref, mmq_q4k_id_run};
let nslots = t * topk;
let (xq, xs, xsum, wqs, wsc, wd, wdm) = gen_mmq_q4k(nslots, n_experts * n, k);
let ids: Vec<u32> = (0..t)
.flat_map(|tok| (0..topk).map(move |j| ((tok * tok + j) % n_experts) as u32))
.collect();
let want = mmq_q4k_id_ref(&xq, &xs, &xsum, &wqs, &wsc, &wd, &wdm, &ids, nslots, n, k);
let (got, ms) = mmq_q4k_id_run::<R>(
client,
&xq,
&xs,
&xsum,
&wqs,
&wsc,
&wd,
&wdm,
&ids,
nslots,
n_experts,
t,
n,
k,
((nslots / n_experts).div_ceil(32)).max(1),
50,
);
let rel = maxabs_over_max(&got, &want);
let gflops = 2.0 * nslots as f64 * n as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] MMQ-Q4K-ID t={} topk={} E={} {}x{} rel={:.2e} {} {:.3} ms {:.0} GFLOP/s",
name,
t,
topk,
n_experts,
n,
k,
rel,
if rel < 1e-2 {
"COOPMAT ✓"
} else {
"MISMATCH ✗"
},
ms,
gflops
);
}
fn maxabs_over_max(got: &[f32], want: &[f32]) -> f32 {
let mut d = 0f32;
let mut r = 1e-9f32;
for (g, w) in got.iter().zip(want) {
d = d.max((g - w).abs());
r = r.max(w.abs());
}
d / r
}
fn gen(rows: usize, k: usize) -> (Vec<f32>, Vec<i32>, Vec<f32>) {
let nb = k / QK8_0;
let mut s = 0x2545F491_4F6CDD1Du64; let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let wd: Vec<f32> = (0..rows * nb)
.map(|_| (next() % 1000) as f32 / 8000.0 + 0.01)
.collect();
let wq: Vec<i32> = (0..rows * k).map(|_| (next() % 255) as i32 - 127).collect();
let x: Vec<f32> = (0..k)
.map(|_| (next() % 2000) as f32 / 1000.0 - 1.0)
.collect();
(wd, wq, x)
}
fn max_rel(a: &[f32], b: &[f32]) -> f32 {
let mut m = 0f32;
for (x, y) in a.iter().zip(b.iter()) {
let d = (x - y).abs();
let denom = x.abs().max(1e-6);
m = m.max(d / denom);
}
m
}
fn check_q8_0_packed<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
rows: usize,
k: usize,
nt: usize,
) {
let (w, x) = gen_q8_0_packed(rows, k);
let reference = matvec_q8_0_packed_ref(&w, &x, rows, k);
let (out, ms) = matvec_q8_0_packed_run::<R>(client, &w, &x, rows, k, nt, 50);
let rel = max_rel(&reference, &out);
let wbytes = rows * (k / 32) * 34; let gbps = wbytes as f64 / (ms * 1e6);
let gflops = 2.0 * rows as f64 * k as f64 / (ms * 1e6);
println!(
"[{:<7}] Q8_0pk {}x{} nt={:<3} max_rel={:.2e} {} {:.3} ms {:.0} GB/s {:.0} GFLOP/s",
name,
rows,
k,
nt,
rel,
if rel < 3e-3 {
"BIT-EXACT ✓"
} else {
"MISMATCH ✗"
},
ms,
gbps,
gflops
);
}
fn check_q8_0_packed_sg<R: Runtime>(
name: &str,
client: &ComputeClient<R>,
rows: usize,
k: usize,
nt: usize,
) {
let (w, x) = gen_q8_0_packed(rows, k);
let reference = matvec_q8_0_packed_ref(&w, &x, rows, k);
let (out, ms) = matvec_q8_0_packed_sg_run::<R>(client, &w, &x, rows, k, nt, 50);
let rel = max_rel(&reference, &out);
let wbytes = rows * (k / 32) * 34;
let gbps = wbytes as f64 / (ms * 1e6);
println!(
"[{:<7}] Q8_0sg {}x{} nt={:<3} max_rel={:.2e} {} {:.3} ms {:.0} GB/s",
name,
rows,
k,
nt,
rel,
if rel < 3e-3 {
"BIT-EXACT ✓"
} else {
"MISMATCH ✗ (plane!=nt?)"
},
ms,
gbps
);
}
fn check<R: Runtime>(name: &str, client: &ComputeClient<R>, rows: usize, k: usize) {
let (wd, wq, x) = gen(rows, k);
let reference = matvec_q8_ref(&wd, &wq, &x, rows, k);
let got = matvec_q8_run::<R>(client, &wd, &wq, &x, rows, k);
let rel = max_rel(&reference, &got);
let ok = rel < 3e-3; for _ in 0..2 {
let _ = matvec_q8_run::<R>(client, &wd, &wq, &x, rows, k);
}
let iters = 20;
let t = Instant::now();
for _ in 0..iters {
let _ = matvec_q8_run::<R>(client, &wd, &wq, &x, rows, k);
}
let _ = (t, iters);
let ms = matvec_q8_bench::<R>(client, &wd, &wq, &x, rows, k, 50);
let gbps = (wd.len() * 4 + wq.len() * 4) as f64 / (ms * 1e6); println!(
"[{:<7}] matvec {}x{} max_rel={:.2e} {} {:.3} ms/dispatch {:.0} GB/s (weight BW)",
name,
rows,
k,
rel,
if ok {
"MATCH ✓ (f32-reorder tol)"
} else {
"MISMATCH ✗"
},
ms,
gbps
);
}
fn check_rms<R: Runtime>(name: &str, client: &ComputeClient<R>, rows: usize, n: usize) {
let mut s = 0x1234_5678_9ABC_DEF1u64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let x: Vec<f32> = (0..rows * n)
.map(|_| (next() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let w: Vec<f32> = (0..n)
.map(|_| (next() % 2000) as f32 / 1000.0 - 1.0)
.collect();
let eps = 1e-6f32;
let mut reference = vec![0f32; rows * n];
for r in 0..rows {
let base = r * n;
let mut ss = 0f32;
for i in 0..n {
let v = x[base + i];
ss += v * v;
}
let denom = (ss / n as f32 + eps).sqrt();
for i in 0..n {
reference[base + i] = x[base + i] / denom * w[i];
}
}
let got = rms_norm_run::<R>(client, &x, &w, rows, n, eps);
let rel = max_rel(&reference, &got);
println!(
"[{:<7}] rmsnorm {}x{} max_rel={:.2e} {}",
name,
rows,
n,
rel,
if rel < 3e-3 {
"MATCH ✓"
} else {
"MISMATCH ✗"
}
);
}
fn main() {
let (rows, k) = (4096usize, 4096usize);
let ctrl = 256usize;
#[cfg(all(feature = "vulkan", feature = "spirv-dump"))]
if std::env::args().any(|a| a == "dump") {
use cubecl::wgpu::{WgpuDevice, WgpuRuntime};
let c = WgpuRuntime::client(&WgpuDevice::default());
check_moe_q4k_blk::<WgpuRuntime>("VK/dump", &c, 128, 768, 8, 2048, 64); check_moe_q4k_blk::<WgpuRuntime>("VK/dump", &c, 128, 2048, 8, 768, 32); check_moe_q4k_dp4a_blk::<WgpuRuntime>("VK/dump", &c, 128, 768, 8, 2048, 64); check_moe_q4k_dp4a_blk::<WgpuRuntime>("VK/dump", &c, 128, 2048, 8, 768, 32); check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/dump", &c, 4096, 2048, 64); check_moe_q6k_blk::<WgpuRuntime>("VK/dump", &c, 128, 2048, 8, 768, 32); check_moe_q6k_dp4a_blk::<WgpuRuntime>("VK/dump", &c, 128, 2048, 8, 768, 32); check_moe_route::<WgpuRuntime>("VK/dump", &c, 8, 128, 8, 128); check_sdpa_blk::<WgpuRuntime>("VK/dump", &c, 32, 8, 1, 2048, 2048, 128, false, 64); check_gemv::<WgpuRuntime>("VK/dump", &c, 128, 4096, 128); check_mmq_q4k::<WgpuRuntime>("VK/dump", &c, 32, 2048, 2048); check_mmq_q4k_rt::<WgpuRuntime>("VK/dump", &c, 32, 2048, 2048); check_mmq_q4k_id::<WgpuRuntime>("VK/dump", &c, 32, 8, 16, 768, 2048);
return;
}
#[cfg(all(feature = "vulkan", feature = "spirv-dump"))]
if std::env::args().any(|a| a == "dumpf32") {
use cubecl::wgpu::{WgpuDevice, WgpuRuntime};
let c = WgpuRuntime::client(&WgpuDevice::default());
let (wqs, wsc, wd, wdm, x) = gen_q4k(2048, 2048);
let _ =
matvec_q4k_f32_blk_run::<WgpuRuntime>(&c, &wqs, &wsc, &wd, &wdm, &x, 2048, 2048, 64, 1);
println!("dumpf32: matvec_q4k_f32_blk nt=64 nr=1 dispatched -- .spv emitted");
return;
}
#[cfg(all(feature = "vulkan", feature = "spirv-dump"))]
if std::env::args().any(|a| a == "dump-flash") {
use cubecl::wgpu::{WgpuDevice, WgpuRuntime};
let r = std::panic::catch_unwind(|| {
let c = WgpuRuntime::client(&WgpuDevice::default());
let q = vec![0.1f32; 2 * 16 * 128];
let k = vec![0.1f32; 1 * 32 * 128];
let v = k.clone();
hanzo_kernel::flash::flash_attn_run::<WgpuRuntime>(
&c, &q, &k, &v, 1, 2, 1, 16, 32, 32, 128, true,
)
});
match r {
Ok(_) => println!("flash dump: dispatch succeeded (coopmat adapter) -- .spv emitted"),
Err(_) => println!("flash dump: dispatch failed (expected on a non-coopmat adapter) -- .spv should already be emitted at codegen"),
}
return;
}
#[cfg(feature = "vulkan")]
if std::env::args().any(|a| a == "coldsweep") {
use cubecl::wgpu::{WgpuDevice, WgpuRuntime};
let c = WgpuRuntime::client(&WgpuDevice::default());
println!("== cache-cold nt sweep :: dense Q4_K decode matvec, footprint >> 32 MB MALL ==");
for (nout, k) in [(16384usize, 8192usize), (16384, 16384)] {
let mb = nout * (k / 256) * 144 / (1024 * 1024);
println!("-- {nout}x{k} ({mb} MB weight footprint) --");
for nt in [32usize, 64, 128, 256] {
check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/cold", &c, nout, k, nt);
}
}
return;
}
#[cfg(feature = "vulkan")]
if std::env::args().any(|a| a == "f32sweep") {
use cubecl::wgpu::{WgpuDevice, WgpuRuntime};
use hanzo_kernel::quant::{matvec_q4k_f32_space, MatvecQ4kF32Eval};
use hanzo_kernel::tune::Evaluator;
let c = WgpuRuntime::client(&WgpuDevice::default());
println!("== cold f32-direct Q4_K matvec sweep (DSL matvec_q4k_f32_blk), bank-rotated >> 32 MB MALL ==");
let space = matvec_q4k_f32_space();
for (nout, k) in [(6144usize, 2048usize), (2048, 2048), (1024, 2048)] {
let wbytes = nout * (k / 256) * 144;
let nbanks = (64 * 1024 * 1024 / wbytes + 1).max(6); let (wqs, wsc, wd, wdm, x) = gen_q4k(nout, k);
let eval =
MatvecQ4kF32Eval::new(&c, &space, &wqs, &wsc, &wd, &wdm, &x, nout, k, nbanks, 6);
println!(
"-- {nout}x{k} ({} MB/bank x {nbanks} banks = {} MB cold footprint) --",
wbytes / (1024 * 1024),
nbanks * wbytes / (1024 * 1024)
);
let mut best = (f64::INFINITY, 0usize, 0usize, 0f64);
for cfg in space.enumerate() {
let nt = cfg.get(&space, "WG") as usize;
let nr = cfg.get(&space, "NR") as usize;
let ms = eval.measure(&cfg, nbanks); if !ms.is_finite() {
println!(" nt={nt:<3} nr={nr} REJECT (diverged from Q4_K oracle)");
continue;
}
let gbps = wbytes as f64 / (ms * 1e6);
println!(" nt={nt:<3} nr={nr} {ms:.4} ms {gbps:.0} GB/s");
if ms < best.0 {
best = (ms, nt, nr, gbps);
}
}
println!(
" >> BEST f32 {nout}x{k}: nt={} nr={} {:.4} ms {:.0} GB/s (worst_rel {:.2e})",
best.1,
best.2,
best.0,
best.3,
eval.worst_rel()
);
use cubecl::server::Handle;
let packed = pack_q4k(&wqs, &wsc, &wd, &wdm);
let (xq, xs, xsum) = quant_act_q8_cpu(&x, 1, k);
let dbanks: Vec<Handle> = (0..nbanks)
.map(|_| c.create_from_slice(u32::as_bytes(&packed)))
.collect();
let xqh = c.create_from_slice(u32::as_bytes(&xq));
let xsh = c.create_from_slice(f32::as_bytes(&xs));
let xsumh = c.create_from_slice(f32::as_bytes(&xsum));
let dmeta = c.create_from_slice(u32::as_bytes(&[k as u32]));
let doh = c.create_from_slice(f32::as_bytes(&vec![0.0f32; nout]));
let dp4a = |nt: usize, bank: &Handle| unsafe {
matvec_q4k_dp4a_blk::launch_unchecked::<f32, WgpuRuntime>(
&c,
Grid::Static(nout as u32, 1, 1),
Block::new_1d(nt as u32),
ArrayArg::from_raw_parts(bank.clone(), packed.len()),
ArrayArg::from_raw_parts(xqh.clone(), xq.len()),
ArrayArg::from_raw_parts(xsh.clone(), xs.len()),
ArrayArg::from_raw_parts(xsumh.clone(), xsum.len()),
ArrayArg::from_raw_parts(doh.clone(), nout),
ArrayArg::from_raw_parts(dmeta.clone(), 1),
nt,
);
};
let mut dbest = (f64::INFINITY, 0usize);
for nt in [32usize, 64, 128] {
for i in 0..(2 * nbanks) {
dp4a(nt, &dbanks[i % nbanks]);
}
let _ = c.read_one_unchecked(doh.clone());
let mut m = f64::INFINITY;
for _ in 0..6 {
let t = std::time::Instant::now();
for i in 0..nbanks {
dp4a(nt, &dbanks[i % nbanks]);
}
let _ = c.read_one_unchecked(doh.clone());
m = m.min(t.elapsed().as_secs_f64() * 1e3 / nbanks as f64);
}
let g = wbytes as f64 / (m * 1e6);
println!(" dp4a nt={nt:<3} {m:.4} ms {g:.0} GB/s");
if m < dbest.0 {
dbest = (m, nt);
}
}
let dg = wbytes as f64 / (dbest.0 * 1e6);
println!(
" >> dp4a BEST {nout}x{k}: nt={} {:.4} ms {:.0} GB/s || f32/dp4a = {:.2}x (needs >= 1.15x to match hand CM)",
dbest.1, dbest.0, dg, best.3 / dg
);
}
return;
}
println!(
"hanzo-kernel :: one #[device] matvec_q8 source, lowered per backend, gated bit-exact\n"
);
#[cfg(feature = "cpu")]
{
use cubecl::cpu::{CpuDevice, CpuRuntime};
let c = CpuRuntime::client(&CpuDevice::default());
check::<CpuRuntime>("CPU", &c, rows, k);
check::<CpuRuntime>("CPU/ctrl", &c, rows, ctrl);
check_q4k::<CpuRuntime>("CPU", &c, rows, k);
check_moe_q4k::<CpuRuntime>("CPU", &c, 8, 64, 4, 512); check_moe_q6k::<CpuRuntime>("CPU", &c, 8, 64, 4, 512);
check_dp4a::<CpuRuntime>("CPU", &c, rows, k, false); check_rms::<CpuRuntime>("CPU", &c, rows, k);
}
#[cfg(feature = "vulkan")]
{
use cubecl::wgpu::{WgpuDevice, WgpuRuntime};
let c = WgpuRuntime::client(&WgpuDevice::default());
check::<WgpuRuntime>("VULKAN", &c, rows, k);
check::<WgpuRuntime>("VK/ctrl", &c, rows, ctrl);
check_q4k::<WgpuRuntime>("VULKAN", &c, rows, k);
check_mmq_q4k::<WgpuRuntime>("VULKAN", &c, 512, 4096, 4096);
check_mmq_q4k_rt::<WgpuRuntime>("VULKAN", &c, 512, 4096, 4096);
check_mmq_q4k_rt::<WgpuRuntime>("VK/tail", &c, 100, 4032, 4096);
check_moe_q4k::<WgpuRuntime>("VULKAN", &c, 128, 768, 8, 2048);
check_moe_q4k_blk::<WgpuRuntime>("VK/blk", &c, 128, 768, 8, 2048, 32); check_moe_q4k_blk::<WgpuRuntime>("VK/blk", &c, 128, 768, 8, 2048, 64);
check_moe_q4k_blk::<WgpuRuntime>("VK/down", &c, 128, 2048, 8, 768, 8);
check_moe_q4k_blk::<WgpuRuntime>("VK/down", &c, 128, 2048, 8, 768, 16);
check_moe_q4k_blk::<WgpuRuntime>("VK/down", &c, 128, 2048, 8, 768, 32);
check_moe_q4k_dp4a_blk::<WgpuRuntime>("VK/dp4a", &c, 128, 768, 8, 2048, 64); check_moe_q4k_dp4a_blk::<WgpuRuntime>("VK/dp4a", &c, 128, 768, 8, 2048, 32);
check_moe_q4k_dp4a_blk::<WgpuRuntime>("VK/dp4a", &c, 128, 2048, 8, 768, 32); check_moe_q4k_dp4a_blk::<WgpuRuntime>("VK/dp4a", &c, 128, 2048, 8, 768, 16);
check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/dproj", &c, 4096, 2048, 64);
check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/dproj", &c, 4096, 2048, 32);
check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/dproj", &c, 2048, 4096, 64);
check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/dproj", &c, 2048, 4096, 128);
check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/dproj", &c, 512, 2048, 64);
check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/cold", &c, 32768, 2048, 32);
check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/cold", &c, 32768, 2048, 64);
check_matvec_q4k_dp4a_blk::<WgpuRuntime>("VK/cold", &c, 32768, 2048, 128);
check_moe_q6k::<WgpuRuntime>("VULKAN", &c, 128, 768, 8, 2048); check_moe_q6k_blk::<WgpuRuntime>("VK/blk", &c, 128, 768, 8, 2048, 32);
check_moe_q6k_blk::<WgpuRuntime>("VK/blk", &c, 128, 768, 8, 2048, 64);
check_moe_q6k_blk::<WgpuRuntime>("VK/down", &c, 128, 2048, 8, 768, 8);
check_moe_q6k_blk::<WgpuRuntime>("VK/down", &c, 128, 2048, 8, 768, 16);
check_moe_q6k_blk::<WgpuRuntime>("VK/down", &c, 128, 2048, 8, 768, 32);
check_moe_route::<WgpuRuntime>("VULKAN", &c, 8, 128, 8, 128);
check_moe_route::<WgpuRuntime>("VK/pref", &c, 512, 128, 8, 128);
check_sdpa_blk::<WgpuRuntime>("VK/attn", &c, 32, 8, 1, 2048, 2048, 128, false, 32); check_sdpa_blk::<WgpuRuntime>("VK/attn", &c, 32, 8, 1, 2048, 2048, 128, false, 64); check_sdpa_blk::<WgpuRuntime>("VK/strd", &c, 32, 8, 1, 2048, 4096, 128, false, 64); check_sdpa_blk::<WgpuRuntime>("VK/attn", &c, 32, 8, 1, 4096, 4096, 128, false, 64); check_sdpa_blk::<WgpuRuntime>("VK/pref", &c, 32, 8, 128, 128, 128, 128, true, 64);
check_gemv::<WgpuRuntime>("VK/gemv", &c, 128, 4096, 32);
check_gemv::<WgpuRuntime>("VK/gemv", &c, 128, 4096, 64);
check_gemv::<WgpuRuntime>("VK/gemv", &c, 128, 4096, 128);
check_gemv::<WgpuRuntime>("VK/gemv", &c, 128, 4096, 256);
check_dp4a::<WgpuRuntime>("VULKAN", &c, rows, k, true);
check_dp4a::<WgpuRuntime>("VK/big", &c, 8192, 8192, true); check_q8_0_packed::<WgpuRuntime>("VULKAN", &c, rows, k, 64);
check_q8_0_packed::<WgpuRuntime>("VK/big", &c, 8192, 8192, 128); check_q8_0_packed_sg::<WgpuRuntime>("VK/sg32", &c, rows, k, 32);
check_q8_0_packed_sg::<WgpuRuntime>("VK/sg64", &c, rows, k, 64);
check_q8_0_packed_sg::<WgpuRuntime>("VK/sgBIG", &c, 8192, 8192, 32);
}
#[cfg(feature = "metal")]
{
use cubecl::wgpu::{WgpuDevice, WgpuRuntime};
let c = WgpuRuntime::client(&WgpuDevice::default());
check::<WgpuRuntime>("METAL", &c, rows, k);
check_q4k::<WgpuRuntime>("METAL", &c, rows, k);
check_moe_q4k::<WgpuRuntime>("METAL", &c, 128, 768, 8, 2048);
check_dp4a::<WgpuRuntime>("METAL", &c, rows, k, true);
}
#[cfg(feature = "cuda")]
{
use cubecl::cuda::{CudaDevice, CudaRuntime};
let c = CudaRuntime::client(&CudaDevice::default());
check::<CudaRuntime>("CUDA", &c, rows, k);
check_q4k::<CudaRuntime>("CUDA", &c, rows, k);
check_dp4a::<CudaRuntime>("CUDA", &c, rows, k, true);
check_rms::<CudaRuntime>("CUDA", &c, rows, k);
}
#[cfg(feature = "rocm")]
{
use hanzo_cubecl_hip::{AmdDevice, HipRuntime};
let c = HipRuntime::client(&AmdDevice::default());
check::<HipRuntime>("ROCM", &c, rows, k);
check_q4k::<HipRuntime>("ROCM", &c, rows, k);
check_dp4a::<HipRuntime>("ROCM", &c, rows, k, true);
}
}