use crate::harness::{BENCH_GUARD, fill, measure};
#[cfg(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64"))]
use crate::harness::{NATIVE_LABEL, NATIVE_MR, NATIVE_NR, NativeTok};
#[cfg(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64"))]
use gemmkit::Workspace;
#[cfg(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64"))]
use gemmkit::driver;
#[cfg(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64"))]
use gemmkit::kernel::FloatGemm;
use gemmkit::{MatMut, MatRef, Parallelism, gemm};
#[cfg(not(target_family = "wasm"))]
fn bench_one(s: usize, parallel: bool) {
let (m, k, n) = (s, s, s);
let a = fill(m * k, 1);
let b = fill(k * n, 2);
let mut c = vec![0.0f32; m * n];
let par = if parallel {
Parallelism::Rayon(0)
} else {
Parallelism::Serial
};
let s_kit = measure(m, k, n, || {
gemm(
1.0,
MatRef::from_col_major(&a, m, k),
MatRef::from_col_major(&b, k, n),
0.0,
MatMut::from_col_major(&mut c, m, n),
par,
);
});
let gpar = if parallel {
gemm::Parallelism::Rayon(0)
} else {
gemm::Parallelism::None
};
let s_gemm = measure(m, k, n, || unsafe {
gemm::gemm(
m,
n,
k,
c.as_mut_ptr(),
m as isize,
1,
false,
a.as_ptr(),
m as isize,
1,
b.as_ptr(),
k as isize,
1,
0.0,
1.0,
false,
false,
false,
gpar,
);
});
let mode = if parallel { "par" } else { "ser" };
print!(
" n={s:<5} {mode} gemmkit={:7.1} (±{:>2.0}%) gemm={:7.1} (±{:>2.0}%) ({:.0}% of gemm)",
s_kit.median,
s_kit.spread_pct(),
s_gemm.median,
s_gemm.spread_pct(),
100.0 * s_kit.median / s_gemm.median.max(1e-9)
);
if !parallel {
let s_mm = measure(m, k, n, || unsafe {
matrixmultiply::sgemm(
m,
k,
n,
1.0,
a.as_ptr(),
1,
m as isize,
b.as_ptr(),
1,
k as isize,
0.0,
c.as_mut_ptr(),
1,
m as isize,
);
});
print!(
" mm={:7.1} ({:.2}x mm)",
s_mm.median,
s_kit.median / s_mm.median.max(1e-9)
);
}
println!();
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64"))]
fn bench_native_equal_isa(s: usize) {
let (m, k, n) = (s, s, s);
let a = fill(m * k, 1);
let b = fill(k * n, 2);
let mut c = vec![0.0f32; m * n];
let mut ws = Workspace::new();
let s_kit = measure(m, k, n, || unsafe {
driver::run::<FloatGemm<f32>, NativeTok, NATIVE_MR, NATIVE_NR>(
NativeTok::default(),
m,
k,
n,
1.0,
a.as_ptr(),
1,
m as isize,
b.as_ptr(),
1,
k as isize,
0.0,
c.as_mut_ptr(),
1,
m as isize,
Parallelism::Serial,
&mut ws,
);
});
let s_gemm = measure(m, k, n, || unsafe {
gemm::gemm(
m,
n,
k,
c.as_mut_ptr(),
m as isize,
1,
false,
a.as_ptr(),
m as isize,
1,
b.as_ptr(),
k as isize,
1,
0.0,
1.0,
false,
false,
false,
gemm::Parallelism::None,
);
});
let label = NATIVE_LABEL;
println!(
" n={s:<5} ser gemmkit-{label}={:7.1} (±{:>2.0}%) gemm-{label}={:7.1} ({:.0}% of gemm)",
s_kit.median,
s_kit.spread_pct(),
s_gemm.median,
100.0 * s_kit.median / s_gemm.median
);
}
#[cfg(not(target_family = "wasm"))]
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_sgemm() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!("\nsgemm GFLOP/s (f32, column-major) — gemmkit best-ISA vs gemm default:");
for &s in &[256usize, 512, 1024, 2048] {
bench_one(s, false);
}
for &s in &[512usize, 1024, 2048, 4096] {
bench_one(s, true);
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64"))]
{
println!("\nequal-ISA (gemmkit vs gemm, same single ISA), single-threaded:");
for &s in &[256usize, 512, 1024, 2048] {
bench_native_equal_isa(s);
}
}
}
#[cfg(not(target_family = "wasm"))]
fn bench_call_latency(m: usize, k: usize, n: usize, par: Parallelism) {
let a = fill(m * k, 1);
let b = fill(k * n, 2);
let mut c = vec![0.0f32; m * n];
let st = measure(m, k, n, || {
gemm(
1.0,
MatRef::from_col_major(&a, m, k),
MatRef::from_col_major(&b, k, n),
0.0,
MatMut::from_col_major(&mut c, m, n),
par,
);
});
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
let us = 2.0 * m as f64 * k as f64 * n as f64 / st.median / 1e3;
println!(
" {m:>4}x{k:<4}x{n:<4} {mode} {:7.1} GFLOP/s (±{:>2.0}%) {us:7.2} us/call",
st.median,
st.spread_pct()
);
}
#[cfg(not(target_family = "wasm"))]
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_call_latency() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!("\nper-call latency at the parallel gate (fixed resolve cost visibility):");
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
bench_call_latency(48, 256, 48, par);
bench_call_latency(96, 256, 96, par);
bench_call_latency(128, 256, 128, par);
}
}
#[cfg(not(target_family = "wasm"))]
fn native_default_tile() -> (usize, usize) {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
(32, 12)
}
#[cfg(target_arch = "aarch64")]
{
(16, 4)
}
#[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))]
{
(4, 4)
}
}
#[cfg(not(target_family = "wasm"))]
fn expected_auto_width(m: usize, k: usize, n: usize, avail: usize, n_jobs: usize) -> usize {
let mnk = m * k * n;
let want = mnk / gemmkit::tuning::par_mnk_per_worker().max(1);
let classes = gemmkit::tuning::pool_classes().min(3);
let mut tiers = [0usize; 3];
let mut n_tiers = 0;
if classes > 0 {
let mut div = 1usize << classes;
while div >= 2 {
let size = avail / div;
if size >= 2 && size < avail {
tiers[n_tiers] = size;
n_tiers += 1;
}
div /= 2;
}
}
let w = if n_tiers == 0 {
want
} else {
#[cfg(target_arch = "aarch64")]
const FULL_WIDTH_AUTO: usize = 14_000_000;
#[cfg(not(target_arch = "aarch64"))]
const FULL_WIDTH_AUTO: usize = 110_000_000;
let full = match gemmkit::tuning::full_width_mnk() {
0 => FULL_WIDTH_AUTO,
v => v,
};
if mnk >= full {
avail
} else {
tiers[..n_tiers]
.iter()
.copied()
.find(|&t| want <= (3 * t) / 2)
.unwrap_or(tiers[n_tiers - 1])
}
};
w.min(avail).min(n_jobs).max(1)
}
#[cfg(not(target_family = "wasm"))]
fn bench_scaling(s: usize) {
let (m, k, n) = (s, s, s);
let a = fill(m * k, 1);
let b = fill(k * n, 2);
let mut c = vec![0.0f32; m * n];
let (mr, nr) = native_default_tile();
let blk = gemmkit::topology().blocking(mr, nr, 4, m, n, k);
let mc = blk.mc.next_multiple_of(mr).max(mr);
let nc = blk.nc.next_multiple_of(nr).max(nr);
let n_jobs = m.div_ceil(mc) * n.min(nc).div_ceil(nr);
println!(
"\n n={s} kc={} mc={} nc={} ~{} jobs/region (tile {mr}x{nr}):",
blk.kc, mc, nc, n_jobs
);
println!(" thr | gemmkit spd eff% | spread | gemm spd");
let base = measure(m, k, n, || {
gemm(
1.0,
MatRef::from_col_major(&a, m, k),
MatRef::from_col_major(&b, k, n),
0.0,
MatMut::from_col_major(&mut c, m, n),
Parallelism::Serial,
);
});
let gbase = measure(m, k, n, || unsafe {
gemm::gemm(
m,
n,
k,
c.as_mut_ptr(),
m as isize,
1,
false,
a.as_ptr(),
m as isize,
1,
b.as_ptr(),
k as isize,
1,
0.0,
1.0,
false,
false,
false,
gemm::Parallelism::None,
);
});
println!(
" 1 | {:9.1} 1.0x 100% | {:5.0}% | {:8.1} 1.0x",
base.median,
base.spread_pct(),
gbase.median
);
let avail = std::thread::available_parallelism()
.map(|x| x.get())
.unwrap_or(1);
for &t in &[2usize, 4, 8, 16, 32] {
let sk = measure(m, k, n, || {
gemm(
1.0,
MatRef::from_col_major(&a, m, k),
MatRef::from_col_major(&b, k, n),
0.0,
MatMut::from_col_major(&mut c, m, n),
Parallelism::Rayon(t),
);
});
let sg = measure(m, k, n, || unsafe {
gemm::gemm(
m,
n,
k,
c.as_mut_ptr(),
m as isize,
1,
false,
a.as_ptr(),
m as isize,
1,
b.as_ptr(),
k as isize,
1,
0.0,
1.0,
false,
false,
false,
gemm::Parallelism::Rayon(t),
);
});
let spd = sk.median / base.median.max(1e-9);
let workers = t.min(avail).min(n_jobs).max(1);
println!(
" {t:3} | {:9.1} {:4.1}x {:3.0}% | {:5.0}% | {:8.1} {:4.1}x",
sk.median,
spd,
100.0 * spd / workers as f64,
sk.spread_pct(),
sg.median,
sg.median / gbase.median.max(1e-9)
);
}
let auto_w = expected_auto_width(m, k, n, avail, n_jobs);
let sk = measure(m, k, n, || {
gemm(
1.0,
MatRef::from_col_major(&a, m, k),
MatRef::from_col_major(&b, k, n),
0.0,
MatMut::from_col_major(&mut c, m, n),
Parallelism::Rayon(0),
);
});
let spd = sk.median / base.median.max(1e-9);
println!(
" auto | {:9.1} {:4.1}x {:3.0}% | {:5.0}% | picks {auto_w} workers",
sk.median,
spd,
100.0 * spd / auto_w as f64,
sk.spread_pct()
);
}
#[cfg(not(target_family = "wasm"))]
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_scaling() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!("\nparallel thread-scaling (f32 col-major) — gemmkit default ISA vs gemm:");
for &s in &[256usize, 512, 1024, 2048] {
bench_scaling(s);
}
}