use anyhow::{Context, ensure};
use cera::backend::wgpu::{DevicePollExt, GpuContext, shaders};
use cera::tensor::DType;
fn env_u32(key: &str, default: u32) -> anyhow::Result<u32> {
let Some(raw) = std::env::var_os(key) else {
return Ok(default);
};
let raw = raw.to_string_lossy();
let value: u32 = raw
.trim()
.parse()
.with_context(|| format!("{key}={raw:?} is not a u32"))?;
ensure!(value > 0, "{key} must be at least 1, got 0");
Ok(value)
}
fn env_f64(key: &str, default: f64) -> anyhow::Result<f64> {
let Some(raw) = std::env::var_os(key) else {
return Ok(default);
};
let raw = raw.to_string_lossy();
let value: f64 = raw
.trim()
.parse()
.with_context(|| format!("{key}={raw:?} is not an f64"))?;
ensure!(
value.is_finite() && value > 0.0,
"{key} must be a finite positive number, got {value}"
);
Ok(value)
}
struct Kernel {
name: &'static str,
dtype: DType,
shader: &'static str,
entry: &'static str,
rows_per_wg: u32,
}
fn block_geom(dtype: DType) -> (usize, usize) {
match dtype {
DType::Q4_0 => (32, 18),
DType::Q8_0 => (32, 34),
DType::Q4KM => (256, 144),
DType::Q5KM => (256, 176),
DType::Q6K => (256, 210),
DType::F32 => (1, 4),
other => panic!("unsupported dtype {other:?}"),
}
}
fn synth_weights(dtype: DType, m: usize, k: usize) -> Vec<u8> {
let (elems, bytes) = block_geom(dtype);
assert_eq!(k % elems, 0, "k must be a multiple of the block size");
let total = m * (k / elems) * bytes;
let mut out: Vec<u8> = (0..total).map(|i| 0x30u8 + (i * 7 % 8) as u8).collect();
out.resize(total.next_multiple_of(4), 0);
out
}
fn main() -> anyhow::Result<()> {
let ctx = GpuContext::new()?;
eprintln!("adapter: {} ({})", ctx.adapter_name, ctx.backend);
let iters = env_u32("ITERS", 200)?;
let iters_2n = iters
.checked_mul(2)
.context("ITERS is too large: the second timing point runs 2x ITERS")?;
let custom = if std::env::var_os("M").is_some() || std::env::var_os("K").is_some() {
let (m, k) = (env_u32("M", 2816)?, env_u32("K", 1024)?);
ensure!(
k.is_multiple_of(256),
"K must be a multiple of 256 (the Q4_K/Q5_K/Q6_K super-block), got {k}"
);
Some(vec![(m, k, "custom")])
} else {
None
};
let default_shapes: &[(u32, u32, &str)] = &[
(2816, 1024, "ffn gate/up"),
(1024, 2816, "ffn down"),
(1024, 1024, "attn qkv/out"),
(65536, 1024, "lm head"),
];
let shapes: &[(u32, u32, &str)] = custom.as_deref().unwrap_or(default_shapes);
let kernels = &[
Kernel {
name: "q4_k",
dtype: DType::Q4KM,
shader: shaders::GEMV_Q4_K,
entry: "gemv_q4_k",
rows_per_wg: 2,
},
Kernel {
name: "q5_k",
dtype: DType::Q5KM,
shader: shaders::GEMV_Q5_K,
entry: "gemv_q5_k",
rows_per_wg: 2,
},
Kernel {
name: "q4_0",
dtype: DType::Q4_0,
shader: shaders::GEMV_Q4_0_FAST,
entry: "gemv_q4_0_fast",
rows_per_wg: 4,
},
Kernel {
name: "q8_0",
dtype: DType::Q8_0,
shader: shaders::GEMV_Q8_0,
entry: "gemv_q8_0",
rows_per_wg: 8,
},
Kernel {
name: "f32",
dtype: DType::F32,
shader: shaders::GEMV_F32,
entry: "gemv_f32",
rows_per_wg: 8,
},
Kernel {
name: "q6_k",
dtype: DType::Q6K,
shader: shaders::GEMV_Q6_K,
entry: "gemv_q6_k",
rows_per_wg: 2,
},
];
println!(
"\n{:<6} {:<14} {:>6} {:>6} {:>10} {:>10} {:>9}",
"kernel", "shape", "m", "k", "ms", "GB/s", "% peak*"
);
println!("{}", "-".repeat(70));
let peak_gbs = env_f64("PEAK_GBS", 400.0)?;
let pipelines: Vec<wgpu::ComputePipeline> = kernels
.iter()
.map(|kern| ctx.create_pipeline(kern.shader, kern.entry, kern.entry))
.collect();
for round in 0..2 {
for (kern, pipeline) in kernels.iter().zip(&pipelines) {
for &(m, k, label) in shapes {
let weight_bytes = synth_weights(kern.dtype, m as usize, k as usize);
let x: Vec<f32> = (0..k)
.map(|i| ((i * 13 + 7) % 251) as f32 * 0.01 - 1.25)
.collect();
let a_buf = ctx.upload_storage(&weight_bytes, "A");
let x_buf = ctx.upload_f32(&x, "x");
let y_buf = ctx.create_storage_rw(u64::from(m) * 4, "y");
let params = [m, k, 0u32, 0u32];
let params_buf = ctx.upload_storage(bytemuck::cast_slice(¶ms), "params");
let bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: a_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let groups = m.div_ceil(kern.rows_per_wg);
let run = |count: u32| {
let mut enc = ctx.device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(pipeline);
pass.set_bind_group(0, &bg, &[]);
for _ in 0..count {
pass.dispatch_workgroups(groups, 1, 1);
}
}
ctx.queue.submit(Some(enc.finish()));
ctx.device.poll_wait();
};
run(2);
let t_n = {
let t0 = std::time::Instant::now();
run(iters);
t0.elapsed()
};
let t_2n = {
let t0 = std::time::Instant::now();
run(iters_2n);
t0.elapsed()
};
let per_iter = t_2n.saturating_sub(t_n).as_secs_f64() / f64::from(iters);
if round == 0 {
continue;
}
let resolvable = t_2n.as_secs_f64() >= t_n.as_secs_f64() * 1.2;
if per_iter <= 0.0 || !resolvable {
println!(
"{:<6} {:<14} {:>6} {:>6} {:>10} {:>10} {:>9}",
kern.name, label, m, k, "noise", "noise", "noise"
);
continue;
}
let (be, bb) = block_geom(kern.dtype);
let bytes = f64::from(m) * f64::from(k) * (bb as f64 / be as f64);
let gbs = bytes / per_iter / 1e9;
println!(
"{:<6} {:<14} {:>6} {:>6} {:>10.4} {:>10.1} {:>8.1}%",
kern.name,
label,
m,
k,
per_iter * 1e3,
gbs,
100.0 * gbs / peak_gbs,
);
}
}
}
println!("\n* % of {peak_gbs} GB/s (override with PEAK_GBS)");
Ok(())
}