use cera::backend::metal::{MetalContext, shaders};
use metal::MTLSize;
const ITERS: u64 = 200;
const ROUNDS: usize = 7;
const N: usize = 1 << 20;
fn size1d(w: u64) -> MTLSize {
MTLSize {
width: w,
height: 1,
depth: 1,
}
}
fn time_dispatches(
ctx: &MetalContext,
pipeline: &metal::ComputePipelineState,
bufs: &[&metal::Buffer],
groups: u64,
) -> f64 {
let start = std::time::Instant::now();
let cb = ctx.queue.new_command_buffer();
let enc = cb.new_compute_command_encoder();
enc.set_compute_pipeline_state(pipeline);
for (i, b) in bufs.iter().enumerate() {
enc.set_buffer(i as u64, Some(b), 0);
}
for _ in 0..ITERS {
enc.dispatch_thread_groups(size1d(groups), size1d(256));
}
enc.end_encoding();
cb.commit();
cb.wait_until_completed();
start.elapsed().as_secs_f64() * 1e6 / ITERS as f64
}
fn run_once(
ctx: &MetalContext,
pipeline: &metal::ComputePipelineState,
bufs: &[&metal::Buffer],
groups: u64,
) {
let cb = ctx.queue.new_command_buffer();
let enc = cb.new_compute_command_encoder();
enc.set_compute_pipeline_state(pipeline);
for (i, b) in bufs.iter().enumerate() {
enc.set_buffer(i as u64, Some(b), 0);
}
enc.dispatch_thread_groups(size1d(groups), size1d(256));
enc.end_encoding();
cb.commit();
cb.wait_until_completed();
}
fn median(mut v: Vec<f64>) -> f64 {
v.sort_by(|a, b| a.partial_cmp(b).expect("no NaN timings"));
v[v.len() / 2]
}
fn rel_diff(a: &[f32], b: &[f32]) -> (f32, f32) {
let max_abs = a
.iter()
.zip(b)
.map(|(x, y)| (x - y).abs())
.fold(0.0f32, f32::max);
let max_ref = a.iter().fold(0.0f32, |m, v| m.max(v.abs()));
(max_abs, max_ref)
}
fn compare(
ctx: &MetalContext,
hand: &metal::ComputePipelineState,
slang: &metal::ComputePipelineState,
hand_bufs: &[&metal::Buffer],
slang_bufs: &[&metal::Buffer],
groups: u64,
) -> (f64, f64) {
time_dispatches(ctx, hand, hand_bufs, groups);
time_dispatches(ctx, slang, slang_bufs, groups);
let mut th = Vec::with_capacity(ROUNDS);
let mut ts = Vec::with_capacity(ROUNDS);
for r in 0..ROUNDS {
if r % 2 == 0 {
th.push(time_dispatches(ctx, hand, hand_bufs, groups));
ts.push(time_dispatches(ctx, slang, slang_bufs, groups));
} else {
ts.push(time_dispatches(ctx, slang, slang_bufs, groups));
th.push(time_dispatches(ctx, hand, hand_bufs, groups));
}
}
(median(th), median(ts))
}
fn print_row(label: &str, h: f64, s: f64) {
println!("{label:<22} {h:>10.2} {s:>10.2} {:>7.3}x", h / s);
}
fn report_agreement(label: &str, a: &[f32], b: &[f32]) -> bool {
let (max_abs, max_ref) = rel_diff(a, b);
let ok = max_abs <= 1e-5 * max_ref.max(1e-6);
println!(
" {label:<20} max_abs_diff={max_abs:.3e} (max|ref|={max_ref:.3e}) {}",
if ok { "MATCH" } else { "MISMATCH" }
);
ok
}
fn main() {
let ctx = match MetalContext::new() {
Ok(c) => c,
Err(e) => {
eprintln!("no Metal device: {e}");
std::process::exit(1);
}
};
let groups = (N as u64).div_ceil(256);
let x: Vec<f32> = (0..N).map(|i| ((i as f32) * 0.001).sin() * 3.0).collect();
let bvals: Vec<f32> = (0..N)
.map(|i| ((i as f32) * 0.0017).cos() * 0.25 + 0.75)
.collect();
println!("== agreement (generated vs handwritten, single dispatch) ==");
let mut ok = true;
let ew_ops: [(&str, u32); 4] = [
("add_inplace", 0),
("mul_inplace", 0),
("scaled_add_inplace", 1.0f32.to_bits()),
("silu_mul_inplace", 0),
];
let ew_hand = shaders::ELEMENTWISE;
let ew_slang = shaders::ELEMENTWISE_SLANG;
for (entry, scale) in ew_ops {
let hand = ctx
.create_pipeline(ew_hand, entry)
.expect("elementwise.metal");
let slang = ctx
.create_pipeline(ew_slang, entry)
.expect("elementwise slang");
let params = ctx.upload_bytes(bytemuck::cast_slice(&[N as u32, scale]));
let bh = ctx.upload_f32(&bvals);
let ah = ctx.upload_f32(&x);
run_once(&ctx, &hand, &[&ah, &bh, ¶ms], groups);
let a = ctx.read_f32(&ah, N);
let bs = ctx.upload_f32(&bvals);
let as_ = ctx.upload_f32(&x);
run_once(&ctx, &slang, &[&as_, &bs, ¶ms], groups);
let b = ctx.read_f32(&as_, N);
ok &= report_agreement(entry, &a, &b);
}
if !ok {
println!("\nkernels disagree; timings below are not comparable");
}
println!(
"\n== timing (us per dispatch, median of {ROUNDS} rounds, {ITERS} dispatches each) =="
);
println!(
"{:<22} {:>10} {:>10} {:>8}",
"kernel", "handwritten", "generated", "ratio"
);
for (entry, scale) in ew_ops {
let hand = ctx
.create_pipeline(ew_hand, entry)
.expect("elementwise.metal");
let slang = ctx
.create_pipeline(ew_slang, entry)
.expect("elementwise slang");
let params = ctx.upload_bytes(bytemuck::cast_slice(&[N as u32, scale]));
let ah = ctx.upload_f32(&x);
let bh = ctx.upload_f32(&bvals);
let as_ = ctx.upload_f32(&x);
let bs = ctx.upload_f32(&bvals);
let (h, s) = compare(
&ctx,
&hand,
&slang,
&[&ah, &bh, ¶ms],
&[&as_, &bs, ¶ms],
groups,
);
print_row(entry, h, s);
}
println!(
"\nratio > 1.00 means the generated kernel is faster; < 1.00 the handwritten one is.\n\
These are branchless maps running in microseconds, so treat anything within a few\n\
percent of 1.00 as no difference."
);
}