use crate::harness::{BENCH_GUARD, Stat, fill, measure_gbps};
use gemmkit::{MatMut, MatRef, Parallelism, gemm};
const STREAM_LEN: usize = 64 * 1024 * 1024;
fn stream_triad_serial() -> Stat {
let n = STREAM_LEN;
let b = fill(n, 1);
let c = fill(n, 2);
let mut a = vec![0.0f32; n];
let alpha = 1.5f32;
measure_gbps(3 * n * 4, || {
for i in 0..n {
a[i] = b[i] + alpha * c[i];
}
std::hint::black_box(a.as_ptr());
})
}
fn stream_copy_serial() -> Stat {
let n = STREAM_LEN;
let src = fill(n, 3);
let mut dst = vec![0.0f32; n];
measure_gbps(2 * n * 4, || {
dst.copy_from_slice(&src);
std::hint::black_box(dst.as_ptr());
})
}
fn stream_triad_parallel(threads: usize) -> Stat {
let n = STREAM_LEN;
let b = fill(n, 1);
let c = fill(n, 2);
let mut a = vec![0.0f32; n];
let alpha = 1.5f32;
let threads = threads.max(1);
let chunk = n.div_ceil(threads);
measure_gbps(3 * n * 4, || {
std::thread::scope(|s| {
for (ai, (bi, ci)) in a
.chunks_mut(chunk)
.zip(b.chunks(chunk).zip(c.chunks(chunk)))
{
s.spawn(move || {
for i in 0..ai.len() {
ai[i] = bi[i] + alpha * ci[i];
}
std::hint::black_box(ai.as_ptr());
});
}
});
std::hint::black_box(a.as_ptr());
})
}
fn stream_triad_parallel_peak(avail: usize) -> (usize, Stat) {
let mut best: Option<(usize, Stat)> = None;
for &t in &[2usize, 4, 8, 16, 32] {
if t > avail {
break;
}
let s = stream_triad_parallel(t);
if best
.as_ref()
.is_none_or(|(_, prev): &(usize, Stat)| s.median > prev.median)
{
best = Some((t, s));
}
}
best.unwrap_or_else(|| (1, stream_triad_serial()))
}
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_stream() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!(
"\nSTREAM bandwidth ceiling (f32, {} MiB arrays):",
STREAM_LEN * 4 / (1024 * 1024)
);
let copy = stream_copy_serial();
let triad = stream_triad_serial();
println!(
" serial copy ={:7.1} GB/s (±{:>2.0}%) triad={:7.1} GB/s (±{:>2.0}%)",
copy.median,
copy.spread_pct(),
triad.median,
triad.spread_pct()
);
let avail = std::thread::available_parallelism()
.map(|x| x.get())
.unwrap_or(1);
for &t in &[2usize, 4, 8, 16, 32] {
if t > avail {
break;
}
let s = stream_triad_parallel(t);
println!(
" {t:3} thr triad={:7.1} GB/s (±{:>2.0}%)",
s.median,
s.spread_pct()
);
}
}
#[cfg(not(target_family = "wasm"))]
fn extern_baselines(
m: usize,
k: usize,
n: usize,
bytes: usize,
par: Parallelism,
) -> (f64, Option<f64>) {
let a = fill(m * k, 1);
let b = fill(k * n, 2);
let mut c = vec![0.0f32; m * n];
let gpar = if matches!(par, Parallelism::Serial) {
gemm::Parallelism::None
} else {
gemm::Parallelism::Rayon(0)
};
let g = measure_gbps(bytes, || 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 mm = matches!(par, Parallelism::Serial).then(|| {
measure_gbps(bytes, || 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,
);
})
.median
});
(g.median, mm)
}
#[allow(unused_variables)]
fn baseline_tail(m: usize, k: usize, n: usize, bytes: usize, par: Parallelism, kit: f64) -> String {
#[cfg(not(target_family = "wasm"))]
{
let (g, mm) = extern_baselines(m, k, n, bytes, par);
let mm_s = mm
.map(|v| format!(" mm={v:6.1} (kit {:.2}×)", kit / v.max(1e-9)))
.unwrap_or_default();
format!(" gemm={g:6.1} (kit {:.2}×){mm_s}", kit / g.max(1e-9))
}
#[cfg(target_family = "wasm")]
{
String::new()
}
}
fn bench_gemv(m: usize, k: usize, par: Parallelism, ceiling: f64) {
let a = fill(m * k, 1);
let x = fill(k, 2);
let mut c = vec![0.0f32; m];
let bytes = (m * k + k + m) * 4;
let st = measure_gbps(bytes, || {
gemm(
1.0,
MatRef::from_col_major(&a, m, k),
MatRef::from_col_major(&x, k, 1),
0.0,
MatMut::from_col_major(&mut c, m, 1),
par,
);
});
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
let tail = baseline_tail(m, k, 1, bytes, par, st.median);
println!(
" m={m:<8} k={k:<5} {mode} kit={:7.1} GB/s (±{:>2.0}%) {:3.0}% ceil{tail}",
st.median,
st.spread_pct(),
100.0 * st.median / ceiling.max(1e-9),
);
}
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_gemv() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!("\ngemv (C[m×1] = A[m×k]·x) — GB/s vs STREAM Triad ceiling:");
let avail = std::thread::available_parallelism()
.map(|x| x.get())
.unwrap_or(1);
let ser_ceiling = stream_triad_serial().median;
let (peak_thr, par_ceiling) = stream_triad_parallel_peak(avail);
let par_ceiling = par_ceiling.median;
println!(
" serial ceiling {ser_ceiling:.1} GB/s; parallel ceiling {par_ceiling:.1} GB/s @ {peak_thr} thr"
);
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
let ceiling = if matches!(par, Parallelism::Serial) {
ser_ceiling
} else {
par_ceiling
};
for &k in &[64usize, 256, 1024] {
for &m in &[1024usize, 8192, 65536] {
bench_gemv(m, k, par, ceiling);
}
}
for &(m, k) in &[
(1_048_576usize, 64usize),
(4_194_304, 16),
(8_388_608, 8),
(16_777_216, 8),
(16_777_216, 16), ] {
bench_gemv(m, k, par, ceiling);
}
}
}
fn bench_gemv_dot(m: usize, k: usize, par: Parallelism, ceiling: f64) {
let a = fill(m * k, 1);
let x = fill(k, 2);
let mut c = vec![0.0f32; m];
let bytes = (m * k + k + m) * 4;
let st = measure_gbps(bytes, || {
gemm(
1.0,
MatRef::from_row_major(&a, m, k),
MatRef::from_col_major(&x, k, 1),
0.0,
MatMut::from_col_major(&mut c, m, 1),
par,
);
});
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
let gflops = st.median * (2.0 * m as f64 * k as f64) / bytes as f64;
let tail = baseline_tail(m, k, 1, bytes, par, st.median);
println!(
" m={m:<6} k={k:<6} {mode} kit={:7.1} GB/s (±{:>2.0}%) {:7.1} GFLOP/s {:3.0}% ceil{tail}",
st.median,
st.spread_pct(),
gflops,
100.0 * st.median / ceiling.max(1e-9),
);
}
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_gemv_dot() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!(
"\ngemv dot-layout (C[m×1] = A[m×k]·x, row-major A -> dot_rows) — GB/s + GFLOP/s vs STREAM ceiling:"
);
let avail = std::thread::available_parallelism()
.map(|x| x.get())
.unwrap_or(1);
let ser_ceiling = stream_triad_serial().median;
let (peak_thr, par_ceiling) = stream_triad_parallel_peak(avail);
let par_ceiling = par_ceiling.median;
println!(
" serial ceiling {ser_ceiling:.1} GB/s; parallel ceiling {par_ceiling:.1} GB/s @ {peak_thr} thr"
);
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
let ceiling = if matches!(par, Parallelism::Serial) {
ser_ceiling
} else {
par_ceiling
};
for &(m, k) in &[
(512usize, 512usize), (2048, 2048), (64, 65536), (8192, 8192), ] {
bench_gemv_dot(m, k, par, ceiling);
}
}
}
#[cfg(all(feature = "half", not(target_family = "wasm")))]
fn bench_gemv_mixed<T: gemmkit::GemmScalar>(
label: &str,
m: usize,
k: usize,
row_major_a: bool,
par: Parallelism,
to_t: impl Fn(f32) -> T,
) {
let a: Vec<T> = fill(m * k, 1).iter().map(|&v| to_t(v)).collect();
let x: Vec<T> = fill(k, 2).iter().map(|&v| to_t(v)).collect();
let mut c = vec![T::ZERO; m];
let bytes = (m * k + k + m) * core::mem::size_of::<T>();
let (alpha, beta) = (to_t(1.0), to_t(0.0));
let mut run = || {
measure_gbps(bytes, || {
let a_ref = if row_major_a {
MatRef::from_row_major(&a, m, k)
} else {
MatRef::from_col_major(&a, m, k)
};
gemm(
alpha,
a_ref,
MatRef::from_col_major(&x, k, 1),
beta,
MatMut::from_col_major(&mut c, m, 1),
par,
);
})
};
let prev = gemmkit::tuning::gemv_threshold();
gemmkit::tuning::set_gemv_threshold(usize::MAX - 1); let gv = run();
gemmkit::tuning::set_gemv_threshold(0); let drv = run();
gemmkit::tuning::set_gemv_threshold(prev);
let gflops = gv.median * (2.0 * m as f64 * k as f64) / bytes as f64;
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
let layout = if row_major_a { "dot " } else { "axpy" };
println!(
" {label} {layout} m={m:<7} k={k:<6} {mode} gemv={:7.1} GB/s (±{:>2.0}%) {:8.1} GFLOP/s driver={:7.1} ({:.2}× driver)",
gv.median,
gv.spread_pct(),
gflops,
drv.median,
gv.median / drv.median.max(1e-9),
);
}
#[cfg(all(feature = "half", not(target_family = "wasm")))]
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_gemv_mixed() {
use gemmkit::{bf16, f16};
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!(
"\nf16/bf16 gemv (C[m×1] = A[m×k]·x) — widen gemv route (dot/axpy) vs general driver:"
);
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
for &row_major_a in &[true, false] {
for &(m, k) in &[(4096usize, 4096usize), (65536, 1024), (1024, 65536)] {
bench_gemv_mixed::<f16>("f16 ", m, k, row_major_a, par, f16::from_f32);
bench_gemv_mixed::<bf16>("bf16", m, k, row_major_a, par, bf16::from_f32);
}
}
}
}
#[cfg(not(target_family = "wasm"))]
fn bench_gemv_paths(m: usize, k: usize, par: Parallelism) {
let a = fill(m * k, 1);
let x = fill(k, 2);
let mut c = vec![0.0f32; m];
let bytes = (m * k + k + m) * 4;
let mut run = || {
measure_gbps(bytes, || {
gemm(
1.0,
MatRef::from_col_major(&a, m, k),
MatRef::from_col_major(&x, k, 1),
0.0,
MatMut::from_col_major(&mut c, m, 1),
par,
)
})
};
let prev = gemmkit::tuning::gemv_threshold();
gemmkit::tuning::set_gemv_threshold(usize::MAX - 1); let axpy = run();
gemmkit::tuning::set_gemv_threshold(0); let driver = run();
gemmkit::tuning::set_gemv_threshold(prev);
let (g, _) = extern_baselines(m, k, 1, bytes, par);
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
println!(
" m={m:<8} k={k:<5} {mode} axpy={:7.1} driver={:7.1} ({:.2}× axpy) gemm={:7.1} ({:.2}× axpy)",
axpy.median,
driver.median,
driver.median / axpy.median.max(1e-9),
g,
g / axpy.median.max(1e-9),
);
}
#[cfg(not(target_family = "wasm"))]
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_gemv_paths() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!("\ngemv path investigation — axpy special path vs general driver vs gemm crate:");
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
for &(m, k) in &[
(4_194_304usize, 16usize),
(4_194_304, 32),
(2_097_152, 48),
(2_097_152, 64),
(1_048_576, 96),
] {
bench_gemv_paths(m, k, par);
}
}
}
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_k_stream() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!("\nK_STREAM_MAX calibration — axpy gemv GB/s vs register-block cap (serial):");
struct Restore(usize);
impl Drop for Restore {
fn drop(&mut self) {
gemmkit::tuning::set_k_stream_max(self.0);
}
}
let _restore = Restore(gemmkit::tuning::k_stream_max());
for &(m, ks) in &[
(300_000usize, &[8usize, 16, 24, 32, 48][..]),
(1_048_576, &[8, 16, 24, 32, 48][..]),
(4_194_304, &[8, 16, 24, 32][..]),
(8_388_608, &[8, 16, 32][..]),
(16_777_216, &[8, 16][..]), ] {
for &k in ks {
let a = fill(m * k, 1);
let x = fill(k, 2);
let bytes = (m * k + k + m) * 4;
let mut row = format!(" m={m:<9} k={k:<3} ");
for &cap in &[32usize, 16, 8] {
gemmkit::tuning::set_k_stream_max(cap);
let mut y = vec![0.0f32; m];
let s = measure_gbps(bytes, || {
gemm(
1.0,
MatRef::from_col_major(&a, m, k),
MatRef::from_col_major(&x, k, 1),
0.0,
MatMut::from_col_major(&mut y, m, 1),
Parallelism::Serial,
);
});
row.push_str(&format!("cap={cap:<2}={:6.1} ", s.median));
}
println!("{row}");
}
}
}
fn bench_gevv(m: usize, n: usize, k: usize, par: Parallelism, ceiling: f64) {
let a = fill(m * k, 1);
let b = fill(k * n, 2);
let mut c = vec![0.0f32; m * n];
let bytes = (m * k + k * n + m * n) * 4;
let st = measure_gbps(bytes, || {
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 tail = baseline_tail(m, k, n, bytes, par, st.median);
println!(
" m={m:<5} n={n:<5} k={k} {mode} kit={:7.1} GB/s (±{:>2.0}%) {:3.0}% ceil{tail}",
st.median,
st.spread_pct(),
100.0 * st.median / ceiling.max(1e-9),
);
}
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_gevv() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!(
"\ngevv / skinny GEMM (C[m×n] = A[m×k]·B[k×n], small k) — GB/s vs STREAM Triad ceiling:"
);
let avail = std::thread::available_parallelism()
.map(|x| x.get())
.unwrap_or(1);
let ser_ceiling = stream_triad_serial().median;
let (peak_thr, par_ceiling) = stream_triad_parallel_peak(avail);
let par_ceiling = par_ceiling.median;
println!(
" serial ceiling {ser_ceiling:.1} GB/s; parallel ceiling {par_ceiling:.1} GB/s @ {peak_thr} thr"
);
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
let ceiling = if matches!(par, Parallelism::Serial) {
ser_ceiling
} else {
par_ceiling
};
for &(m, n) in &[(4096usize, 4096usize), (8192, 2048)] {
for &k in &[1usize, 2, 4] {
bench_gevv(m, n, k, par, ceiling);
}
}
}
}
fn bench_gemv_scaling(m: usize, k: usize, ceiling: f64, avail: usize) {
let a = fill(m * k, 1);
let x = fill(k, 2);
let mut c = vec![0.0f32; m];
let bytes = (m * k + k + m) * 4;
let mut run = |par| {
measure_gbps(bytes, || {
gemm(
1.0,
MatRef::from_col_major(&a, m, k),
MatRef::from_col_major(&x, k, 1),
0.0,
MatMut::from_col_major(&mut c, m, 1),
par,
);
})
};
println!(" m={m} k={k} (A={} MiB):", m * k * 4 / (1024 * 1024));
for &t in &[1usize, 2, 4, 8, 16, 32] {
if t > avail {
break;
}
let par = if t == 1 {
Parallelism::Serial
} else {
Parallelism::Rayon(t)
};
let st = run(par);
println!(
" t={t:<3} {:7.1} GB/s (±{:>2.0}%) {:3.0}% ceil",
st.median,
st.spread_pct(),
100.0 * st.median / ceiling.max(1e-9)
);
}
let st = run(Parallelism::Rayon(0));
println!(
" auto {:7.1} GB/s (±{:>2.0}%) {:3.0}% ceil",
st.median,
st.spread_pct(),
100.0 * st.median / ceiling.max(1e-9)
);
}
fn bench_small_k_crossover(m: usize, n: usize, k: usize, par: Parallelism) {
let a = fill(m * k, 1);
let b = fill(k * n, 2);
let mut c = vec![0.0f32; m * n];
let bytes = (m * k + k * n + m * n) * 4;
let mut run = |v: usize| {
gemmkit::tuning::set_small_k_threshold(v);
measure_gbps(bytes, || {
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 prev = gemmkit::tuning::small_k_threshold();
let on = run(k); let off = run(0); gemmkit::tuning::set_small_k_threshold(prev);
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
println!(
" m={m:<5} n={n:<5} k={k:<3} {mode} small_k={:7.1} (±{:>2.0}%) driver={:7.1} (±{:>2.0}%) (small_k {:.0}% of driver)",
on.median,
on.spread_pct(),
off.median,
off.spread_pct(),
100.0 * on.median / off.median.max(1e-9)
);
}
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_small_k() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!(
"\nsmall-k route crossover (skinny GEMM) — in-place small_k vs register-tiling driver:"
);
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
for &(m, n) in &[(4096usize, 4096usize), (8192, 2048)] {
for &k in &[2usize, 4, 8, 16, 32, 64] {
bench_small_k_crossover(m, n, k, par);
}
}
}
}
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_gemv_scaling() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!("\ngemv thread-scaling (forced Rayon(t) vs auto) — GB/s vs parallel STREAM ceiling:");
let avail = std::thread::available_parallelism()
.map(|x| x.get())
.unwrap_or(1);
let (peak_thr, par_ceiling) = stream_triad_parallel_peak(avail);
let ceiling = par_ceiling.median;
println!(" parallel ceiling {ceiling:.1} GB/s @ {peak_thr} thr; cores={avail}");
bench_gemv_scaling(65536, 1024, ceiling, avail);
bench_gemv_scaling(16_777_216, 8, ceiling, avail);
bench_gemv_scaling(1024, 1024, ceiling, avail);
bench_gemv_scaling(8192, 64, ceiling, avail);
}