use crate::harness::{BENCH_GUARD, fill, measure};
use gemmkit::{MatMut, MatRef, Parallelism, gemm};
#[cfg(not(target_family = "wasm"))]
fn with_route<R>(small_mn: usize, small_k: usize, f: impl FnOnce() -> R) -> R {
let (pm, pk) = (
gemmkit::tuning::small_mn_dim(),
gemmkit::tuning::small_k_threshold(),
);
gemmkit::tuning::set_small_mn_dim(small_mn);
gemmkit::tuning::set_small_k_threshold(small_k);
let r = f();
gemmkit::tuning::set_small_mn_dim(pm);
gemmkit::tuning::set_small_k_threshold(pk);
r
}
#[cfg(not(target_family = "wasm"))]
fn extern_gflops_small(
m: usize,
k: usize,
n: usize,
row_major_a: bool,
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 (a_cs, a_rs) = if row_major_a {
(1isize, k as isize)
} else {
(m as isize, 1isize)
};
let gpar = if matches!(par, Parallelism::Serial) {
gemm::Parallelism::None
} else {
gemm::Parallelism::Rayon(0)
};
let g = measure(m, k, n, || unsafe {
gemm::gemm(
m,
n,
k,
c.as_mut_ptr(),
m as isize,
1, false,
a.as_ptr(),
a_cs,
a_rs,
b.as_ptr(),
k as isize,
1, 0.0,
1.0,
false,
false,
false,
gpar,
);
});
let mm = matches!(par, Parallelism::Serial).then(|| {
measure(m, k, n, || unsafe {
matrixmultiply::sgemm(
m,
k,
n,
1.0,
a.as_ptr(),
a_rs,
a_cs, b.as_ptr(),
1,
k as isize, 0.0,
c.as_mut_ptr(),
1,
m as isize, );
})
.median
});
(g.median, mm)
}
#[cfg(not(target_family = "wasm"))]
fn bench_small_mn(m: usize, n: usize, k: usize, row_major_a: bool, par: Parallelism) {
let a = fill(m * k, 1);
let b = fill(k * n, 2);
let mut c = vec![0.0f32; m * n];
let mut run = || {
measure(m, k, n, || {
let av = if row_major_a {
MatRef::from_row_major(&a, m, k)
} else {
MatRef::from_col_major(&a, m, k)
};
gemm(
1.0,
av,
MatRef::from_col_major(&b, k, n),
0.0,
MatMut::from_col_major(&mut c, m, n),
par,
);
})
};
let horiz = with_route(usize::MAX, 0, &mut run);
let smallk = with_route(0, usize::MAX, &mut run);
let driver = with_route(0, 0, &mut run);
let (g, mm) = extern_gflops_small(m, k, n, row_major_a, par);
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
let mm_s = mm
.map(|v| format!(" mm={v:6.1} ({:.2}×)", horiz.median / v.max(1e-9)))
.unwrap_or_default();
println!(
" {m:>2}×{n:<2} k={k:<5} {mode} horiz={:7.1} small_k={:7.1} driver={:7.1} ({:.2}× h) gemm={:6.1} ({:.2}× h){mm_s}",
horiz.median,
smallk.median,
driver.median,
horiz.median / driver.median.max(1e-9),
g,
horiz.median / g.max(1e-9),
);
}
#[cfg(not(target_family = "wasm"))]
fn bench_small_mn_layouts(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 horiz = with_route(usize::MAX, 0, || {
measure(m, k, n, || {
gemm(
1.0,
MatRef::from_row_major(&a, m, k),
MatRef::from_col_major(&b, k, n),
0.0,
MatMut::from_col_major(&mut c, m, n),
par,
);
})
});
let all_rm = measure(m, k, n, || {
gemm(
1.0,
MatRef::from_row_major(&a, m, k),
MatRef::from_row_major(&b, k, n),
0.0,
MatMut::from_row_major(&mut c, m, n),
par,
);
});
let all_cm = 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"
};
println!(
" {m:>2}x{n:<2} k={k:<6} {mode} horiz={:7.1} all_rm={:7.1} ({:.2}x h) all_cm={:7.1} ({:.2}x h)",
horiz.median,
all_rm.median,
horiz.median / all_rm.median.max(1e-9),
all_cm.median,
horiz.median / all_cm.median.max(1e-9),
);
}
#[cfg(not(target_family = "wasm"))]
fn bench_small_mn_pack_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 mut run = || {
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 pmnk = gemmkit::tuning::small_mn_pack_min_k();
gemmkit::tuning::set_small_mn_pack_min_k(0);
let packed = with_route(usize::MAX, 0, &mut run);
gemmkit::tuning::set_small_mn_pack_min_k(pmnk);
let driver = with_route(0, 0, &mut run);
let smallk = with_route(0, usize::MAX, &mut run);
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
println!(
" {m:>2}x{n:<2} k={k:<6} {mode} packed={:7.1} driver={:7.1} ({:.2}x d) small_k={:7.1}",
packed.median,
driver.median,
packed.median / driver.median.max(1e-9),
smallk.median,
);
}
#[cfg(not(target_family = "wasm"))]
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_small_mn_pack() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!(
"\nsmall-m,n pack tier — view 1: recovery vs the eligible horizontal ceiling (defaults):"
);
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
for &(m, n, k) in &[(4usize, 4usize, 65536usize), (8, 8, 65536), (16, 16, 16384)] {
bench_small_mn_layouts(m, n, k, par);
}
}
println!(
"\nsmall-m,n pack tier — view 2: crossover (packed-horizontal vs driver, ineligible all-col-major):"
);
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
for &(m, n) in &[(4usize, 4usize), (8, 8), (16, 16)] {
for &k in &[32usize, 64, 256, 1024, 4096, 16384] {
bench_small_mn_pack_crossover(m, n, k, par);
}
}
}
}
#[cfg(all(feature = "half", not(target_family = "wasm")))]
fn bench_small_mn_f16(m: usize, n: usize, k: usize, par: Parallelism) {
use gemmkit::f16;
let to16 = |v: &[f32]| v.iter().map(|&x| f16::from_f32(x)).collect::<Vec<_>>();
let a = to16(&fill(m * k, 1));
let b = to16(&fill(k * n, 2));
let mut c = vec![f16::from_f32(0.0); m * n];
let mut run = || {
measure(m, k, n, || {
gemm(
f16::from_f32(1.0),
MatRef::from_row_major(&a, m, k),
MatRef::from_col_major(&b, k, n),
f16::from_f32(0.0),
MatMut::from_col_major(&mut c, m, n),
par,
);
})
};
let horiz = with_route(usize::MAX, 0, &mut run);
let driver = with_route(0, 0, &mut run);
let gpar = if matches!(par, Parallelism::Serial) {
gemm::Parallelism::None
} else {
gemm::Parallelism::Rayon(0)
};
let g = measure(m, k, n, || unsafe {
gemm::gemm(
m,
n,
k,
c.as_mut_ptr(),
m as isize,
1, false,
a.as_ptr(),
1,
k as isize, b.as_ptr(),
k as isize,
1, f16::from_f32(0.0),
f16::from_f32(1.0),
false,
false,
false,
gpar,
);
});
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
println!(
" {m:>2}×{n:<2} k={k:<5} {mode} horiz={:7.1} driver={:7.1} ({:.2}× h) gemm={:6.1} ({:.2}× h)",
horiz.median,
driver.median,
horiz.median / driver.median.max(1e-9),
g.median,
horiz.median / g.median.max(1e-9),
);
}
#[cfg(all(feature = "half", not(target_family = "wasm")))]
fn bench_small_mn_bf16(m: usize, n: usize, k: usize, par: Parallelism) {
use gemmkit::bf16;
let to16 = |v: &[f32]| v.iter().map(|&x| bf16::from_f32(x)).collect::<Vec<_>>();
let a = to16(&fill(m * k, 1));
let b = to16(&fill(k * n, 2));
let mut c = vec![bf16::from_f32(0.0); m * n];
let mut run = || {
measure(m, k, n, || {
gemm(
bf16::from_f32(1.0),
MatRef::from_row_major(&a, m, k),
MatRef::from_col_major(&b, k, n),
bf16::from_f32(0.0),
MatMut::from_col_major(&mut c, m, n),
par,
);
})
};
let horiz = with_route(usize::MAX, 0, &mut run);
let driver = with_route(0, 0, &mut run);
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
println!(
" {m:>2}×{n:<2} k={k:<5} {mode} horiz={:7.1} driver={:7.1} ({:.2}× h)",
horiz.median,
driver.median,
horiz.median / driver.median.max(1e-9),
);
}
#[cfg(all(feature = "int8", not(target_family = "wasm")))]
fn bench_small_mn_i8(m: usize, n: usize, k: usize, par: Parallelism) {
let a: Vec<i8> = (0..m * k).map(|i| (i % 17) as i8 - 8).collect();
let b: Vec<i8> = (0..k * n).map(|i| (i % 13) as i8 - 6).collect();
let mut c = vec![0i32; m * n];
let mut run = || {
measure(m, k, n, || {
gemmkit::gemm_i8(
1,
MatRef::from_row_major(&a, m, k),
MatRef::from_col_major(&b, k, n),
0,
MatMut::from_col_major(&mut c, m, n),
par,
);
})
};
let horiz = with_route(usize::MAX, 0, &mut run);
let smallk = with_route(0, usize::MAX, &mut run);
let driver = with_route(0, 0, &mut run);
let mode = if matches!(par, Parallelism::Serial) {
"ser"
} else {
"par"
};
println!(
" {m:>2}×{n:<2} k={k:<5} {mode} horiz={:7.1} small_k={:7.1} driver={:7.1} ({:.2}× h)",
horiz.median,
smallk.median,
driver.median,
horiz.median / driver.median.max(1e-9),
);
}
#[cfg(not(target_family = "wasm"))]
#[test]
#[ignore = "benchmark; run with --release --ignored --nocapture"]
fn perf_small_mn() {
let _guard = BENCH_GUARD.lock().unwrap_or_else(|e| e.into_inner());
println!(
"\nsmall-m,n horizontal route (C[m×n]=A·B, small m,n, long k) — GFLOP/s, row-major A + col-major B (fast-path layout):"
);
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
for &s in &[2usize, 4, 8, 16, 32] {
for &k in &[64usize, 256, 1024, 4096] {
bench_small_mn(s, s, k, true, par);
}
}
for &(m, n) in &[(2usize, 8usize), (8, 2), (4, 16), (16, 4)] {
for &k in &[256usize, 4096] {
bench_small_mn(m, n, k, true, par);
}
}
}
#[cfg(feature = "int8")]
{
println!("\n i8 -> i32 (widen horizontal dot vs vpdpbusd/widen driver):");
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
for &(m, n, k) in &[
(4usize, 4usize, 65536usize),
(8, 8, 65536),
(16, 16, 65536),
(16, 16, 4096),
] {
bench_small_mn_i8(m, n, k, par);
}
}
}
#[cfg(feature = "half")]
{
println!("\n f16 (f32-accumulate mixed horizontal kernel):");
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
for &s in &[4usize, 8, 16, 32] {
for &k in &[256usize, 4096] {
bench_small_mn_f16(s, s, k, par);
}
}
}
println!("\n bf16 (widen horizontal path vs vdpbf16ps VNNI driver on x86):");
for &par in &[Parallelism::Serial, Parallelism::Rayon(0)] {
for &s in &[4usize, 8, 16, 32] {
for &k in &[256usize, 4096] {
bench_small_mn_bf16(s, s, k, par);
}
}
}
}
}