use memra_engine::Engine;
const N_LAYERS: f64 = 45.0;
const PEAK_TBS: f64 = 1.79;
fn main() -> Result<(), Box<dyn std::error::Error>> {
let reps: usize = std::env::args()
.collect::<Vec<_>>()
.windows(2)
.find(|w| w[0] == "--reps")
.and_then(|w| w[1].parse().ok())
.unwrap_or(300);
let e = Engine::new(0)?;
println!("decode-kernel-census reps={reps} peak={PEAK_TBS} TB/s layers={N_LAYERS}");
qkvg(&e, reps, 4096, 4096, 512, 32)?;
b4(&e, reps, 1024, 4096)?;
plain(&e, reps, 1280, 4096, "shexp down 1280->4096")?;
plain(&e, reps, 4096, 64448, "head 4096->64448 (HEAD_SPLIT half)")?;
q8(&e, reps, 4096, 5152, "q8_0 qkv-equivalent 4096->5152")?;
q8(&e, reps, 4096, 4096, "q8_0 o_proj-equivalent 4096->4096")?;
q8(
&e,
reps,
1280,
4096,
"q8_0 shexp-down-equivalent 1280->4096",
)?;
q8(&e, reps, 4096, 64448, "q8_0 head-equivalent 4096->64448")?;
q8(
&e,
reps,
4096,
20480,
"q8_0 expert gate+up stacked 4096->20480",
)?;
Ok(())
}
fn q8(
e: &Engine,
reps: usize,
in_f: usize,
out_f: usize,
label: &str,
) -> Result<(), Box<dyn std::error::Error>> {
const QK: usize = 32;
let row_bytes = in_f / QK * (QK + 2);
let w = e.alloc_u8(out_f * row_bytes)?;
let aq = e.htod_i8(&vec![1i8; in_f])?;
let ad = e.htod(&vec![0.01f32; 2 * in_f / QK])?;
let mut y = e.htod(&vec![0f32; out_f])?;
let mut run = || -> Result<(), Box<dyn std::error::Error>> {
e.qmatvec_mmvq_into(
&w,
&aq,
&ad,
1,
in_f,
out_f,
memra_engine::QT_Q8_0,
row_bytes,
1.0,
true,
&mut y,
)
};
for _ in 0..30 {
run()?;
}
e.stream().synchronize()?;
let t0 = std::time::Instant::now();
for _ in 0..reps {
run()?;
}
e.stream().synchronize()?;
let us = t0.elapsed().as_secs_f64() * 1e6 / reps as f64;
let per_layer = if out_f > 60000 { 1.0 / N_LAYERS } else { 1.0 };
report(label, us, (out_f * row_bytes) as f64, per_layer);
Ok(())
}
fn report(label: &str, us: f64, bytes: f64, per_token_calls: f64) {
let tbs = bytes / us / 1e6;
println!(
"[{label}] {us:8.1} us/call {tbs:5.2} TB/s ({:4.1}% of peak) {:6.2} ms/token @ {per_token_calls} call(s)/layer",
100.0 * tbs / PEAK_TBS,
us * per_token_calls * N_LAYERS / 1000.0
);
}
fn qkvg(
e: &Engine,
reps: usize,
in_f: usize,
out_q: usize,
out_kv: usize,
out_g: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let rows = out_q + 2 * out_kv + out_g;
let wq = e.alloc_u8(out_q * in_f * 2)?;
let wk = e.alloc_u8(out_kv * in_f * 2)?;
let wv = e.alloc_u8(out_kv * in_f * 2)?;
let wg = e.alloc_u8(out_g * in_f * 2)?;
let x = e.htod(&vec![0.01f32; in_f])?;
let mut yq = e.htod(&vec![0f32; out_q])?;
let mut yk = e.htod(&vec![0f32; out_kv])?;
let mut yv = e.htod(&vec![0f32; out_kv])?;
let mut yg = e.htod(&vec![0f32; out_g])?;
let mut run = || -> Result<(), Box<dyn std::error::Error>> {
e.matvec_bf16_qkvg_into(
&wq, &wk, &wv, &wg, &x, &mut yq, &mut yk, &mut yv, &mut yg, in_f, out_q, out_kv, out_g,
)
};
for _ in 0..30 {
run()?;
}
e.stream().synchronize()?;
let t0 = std::time::Instant::now();
for _ in 0..reps {
run()?;
}
e.stream().synchronize()?;
let us = t0.elapsed().as_secs_f64() * 1e6 / reps as f64;
report(
&format!("qkvg {in_f}->{rows} (per card)"),
us,
(rows * in_f * 2) as f64,
1.0,
);
Ok(())
}
fn b4(
e: &Engine,
reps: usize,
block_cols: usize,
out_f: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let w: Vec<_> = (0..4)
.map(|_| e.alloc_u8(out_f * block_cols * 2))
.collect::<Result<_, _>>()?;
let x = e.htod(&vec![0.01f32; 4 * block_cols])?;
let mut y = e.htod(&vec![0f32; out_f])?;
let mut run = || -> Result<(), Box<dyn std::error::Error>> {
e.matvec_bf16_b4_into([&w[0], &w[1], &w[2], &w[3]], &x, &mut y, block_cols, out_f)
};
for _ in 0..30 {
run()?;
}
e.stream().synchronize()?;
let t0 = std::time::Instant::now();
for _ in 0..reps {
run()?;
}
e.stream().synchronize()?;
let us = t0.elapsed().as_secs_f64() * 1e6 / reps as f64;
report(
&format!("b4 o_proj 4x{block_cols}->{out_f}"),
us,
(4 * out_f * block_cols * 2) as f64,
1.0,
);
Ok(())
}
fn plain(
e: &Engine,
reps: usize,
in_f: usize,
out_f: usize,
label: &str,
) -> Result<(), Box<dyn std::error::Error>> {
let w = e.alloc_u8(in_f * out_f * 2)?;
let x = e.htod(&vec![0.01f32; in_f])?;
let mut y = e.htod(&vec![0f32; out_f])?;
let mut run = || -> Result<(), Box<dyn std::error::Error>> {
e.matvec_bf16_into(&w, &x, &mut y, in_f, out_f)
};
for _ in 0..30 {
run()?;
}
e.stream().synchronize()?;
let t0 = std::time::Instant::now();
for _ in 0..reps {
run()?;
}
e.stream().synchronize()?;
let us = t0.elapsed().as_secs_f64() * 1e6 / reps as f64;
let per_layer = if out_f > 60000 { 1.0 / N_LAYERS } else { 1.0 };
report(label, us, (in_f * out_f * 2) as f64, per_layer);
Ok(())
}