use numrs2::array::Array;
fn naive_gemm_f64(m: usize, k: usize, n: usize, a: &[f64], b: &[f64]) -> Vec<f64> {
let mut c = vec![0.0f64; m * n];
for i in 0..m {
for p in 0..k {
let a_ip = a[i * k + p];
for j in 0..n {
c[i * n + j] += a_ip * b[p * n + j];
}
}
}
c
}
fn naive_gemm_f32(m: usize, k: usize, n: usize, a: &[f32], b: &[f32]) -> Vec<f32> {
let mut c = vec![0.0f32; m * n];
for i in 0..m {
for p in 0..k {
let a_ip = a[i * k + p];
for j in 0..n {
c[i * n + j] += a_ip * b[p * n + j];
}
}
}
c
}
fn naive_gemm_i32(m: usize, k: usize, n: usize, a: &[i32], b: &[i32]) -> Vec<i32> {
let mut c = vec![0i32; m * n];
for i in 0..m {
for p in 0..k {
let a_ip = a[i * k + p];
for j in 0..n {
c[i * n + j] += a_ip * b[p * n + j];
}
}
}
c
}
fn assert_close_rel(got: f64, expected: f64, tol: f64, ctx: impl std::fmt::Display) {
let scale = expected.abs().max(1.0);
assert!(
(got - expected).abs() <= tol * scale,
"{ctx}: got {got}, expected {expected} (relative diff {:.3e}, tol {tol:.0e})",
(got - expected).abs() / scale,
);
}
fn assert_close_rel_f32(got: f32, expected: f32, tol: f32, ctx: impl std::fmt::Display) {
let scale = expected.abs().max(1.0);
assert!(
(got - expected).abs() <= tol * scale,
"{ctx}: got {got}, expected {expected} (relative diff {:.3e}, tol {tol:.0e})",
(got - expected).abs() / scale,
);
}
fn seq_f64(len: usize, scale: f64, offset: f64) -> Vec<f64> {
(0..len).map(|i| (i as f64) * scale + offset).collect()
}
fn seq_f32(len: usize, scale: f32, offset: f32) -> Vec<f32> {
(0..len).map(|i| (i as f32) * scale + offset).collect()
}
fn seq_i32(len: usize, modulus: i32, offset: i32) -> Vec<i32> {
(0..len).map(|i| (i as i32 % modulus) + offset).collect()
}
const SIZE_GRID: &[(usize, usize, usize)] = &[
(1, 1, 1),
(1, 32, 1),
(2, 31, 33),
(31, 32, 33),
(32, 32, 32),
(33, 64, 31),
(63, 64, 65),
(64, 64, 64),
(65, 63, 127),
(127, 128, 2),
(128, 128, 128),
(128, 1, 64),
(64, 127, 65),
(32, 128, 63),
(64, 64, 127),
(64, 64, 128),
(64, 64, 129),
];
#[test]
fn matmul_2d_f64_matches_naive_across_boundary_grid() {
for &(m, k, n) in SIZE_GRID {
let a_data = seq_f64(m * k, 0.125, -3.0);
let b_data = seq_f64(k * n, -0.0625, 1.5);
let expected = naive_gemm_f64(m, k, n, &a_data, &b_data);
let a = Array::from_vec(a_data).reshape(&[m, k]);
let b = Array::from_vec(b_data).reshape(&[k, n]);
let c = a.matmul(&b).expect("matmul should succeed");
assert_eq!(c.shape(), vec![m, n], "(m={m},k={k},n={n}) shape");
for (idx, (got, want)) in c.to_vec().iter().zip(&expected).enumerate() {
assert_close_rel(*got, *want, 1e-9, format!("(m={m},k={k},n={n}) idx={idx}"));
}
}
}
#[test]
fn matmul_2d_f32_matches_naive_across_boundary_grid() {
for &(m, k, n) in SIZE_GRID {
let a_data = seq_f32(m * k, 0.125, -3.0);
let b_data = seq_f32(k * n, -0.0625, 1.5);
let expected = naive_gemm_f32(m, k, n, &a_data, &b_data);
let a = Array::from_vec(a_data).reshape(&[m, k]);
let b = Array::from_vec(b_data).reshape(&[k, n]);
let c = a.matmul(&b).expect("matmul should succeed");
assert_eq!(c.shape(), vec![m, n], "(m={m},k={k},n={n}) shape");
for (idx, (got, want)) in c.to_vec().iter().zip(&expected).enumerate() {
assert_close_rel_f32(*got, *want, 1e-4, format!("(m={m},k={k},n={n}) idx={idx}"));
}
}
}
#[test]
fn matmul_2d_i32_generic_tier_matches_naive_across_boundary_grid() {
for &(m, k, n) in SIZE_GRID {
let a_data = seq_i32(m * k, 7, -3);
let b_data = seq_i32(k * n, 5, -2);
let expected = naive_gemm_i32(m, k, n, &a_data, &b_data);
let a = Array::from_vec(a_data).reshape(&[m, k]);
let b = Array::from_vec(b_data).reshape(&[k, n]);
let c = a.matmul(&b).expect("matmul should succeed");
assert_eq!(c.shape(), vec![m, n], "(m={m},k={k},n={n}) shape");
assert_eq!(c.to_vec(), expected, "(m={m},k={k},n={n})");
}
}
#[test]
fn matmul_2d_zero_k_yields_zero_filled_output() {
let (m, k, n) = (5usize, 0usize, 4usize);
let a = Array::<f64>::from_vec(vec![]).reshape(&[m, k]);
let b = Array::<f64>::from_vec(vec![]).reshape(&[k, n]);
let c = a.matmul(&b).expect("matmul with k=0 should succeed");
assert_eq!(c.shape(), vec![m, n]);
assert_eq!(c.to_vec(), vec![0.0; m * n]);
}
#[test]
fn matmul_2d_zero_m_and_zero_n_yield_empty_outputs() {
let a = Array::<f64>::from_vec(vec![]).reshape(&[0, 5]);
let b = Array::from_vec(seq_f64(20, 1.0, 0.0)).reshape(&[5, 4]);
let c = a.matmul(&b).expect("matmul with m=0 should succeed");
assert_eq!(c.shape(), vec![0, 4]);
assert!(c.to_vec().is_empty());
let a2 = Array::from_vec(seq_f64(20, 1.0, 0.0)).reshape(&[5, 4]);
let b2 = Array::<f64>::from_vec(vec![]).reshape(&[4, 0]);
let c2 = a2.matmul(&b2).expect("matmul with n=0 should succeed");
assert_eq!(c2.shape(), vec![5, 0]);
assert!(c2.to_vec().is_empty());
}
#[test]
fn matmul_2d_zero_k_on_generic_tier_yields_zero_filled_output() {
let (m, k, n) = (3usize, 0usize, 6usize);
let a = Array::<i32>::from_vec(vec![]).reshape(&[m, k]);
let b = Array::<i32>::from_vec(vec![]).reshape(&[k, n]);
let c = a.matmul(&b).expect("matmul with k=0 should succeed");
assert_eq!(c.shape(), vec![m, n]);
assert_eq!(c.to_vec(), vec![0i32; m * n]);
}
#[test]
fn matmul_2d_handles_non_contiguous_operands() {
let (m, k, n) = (33usize, 65usize, 31usize);
let a = Array::from_vec(seq_f64(k * m, 0.25, -2.0))
.reshape(&[k, m])
.transpose_axis(0, 1);
let b = Array::from_vec(seq_f64(n * k, -0.125, 0.75))
.reshape(&[n, k])
.transpose_axis(0, 1);
assert_eq!(a.shape(), vec![m, k]);
assert_eq!(b.shape(), vec![k, n]);
assert!(!a.is_c_contiguous(), "test requires a non-contiguous `a`");
assert!(!b.is_c_contiguous(), "test requires a non-contiguous `b`");
let expected = naive_gemm_f64(m, k, n, &a.to_vec(), &b.to_vec());
let c = a.matmul(&b).expect("matmul should succeed");
assert_eq!(c.shape(), vec![m, n]);
for (idx, (got, want)) in c.to_vec().iter().zip(&expected).enumerate() {
assert_close_rel(*got, *want, 1e-9, format!("both transposed idx={idx}"));
}
let b_contig = Array::from_vec(b.to_vec()).reshape(&[k, n]);
assert!(b_contig.is_c_contiguous());
let c_a_only = a.matmul(&b_contig).expect("matmul should succeed");
for (idx, (got, want)) in c_a_only.to_vec().iter().zip(&expected).enumerate() {
assert_close_rel(*got, *want, 1e-9, format!("a transposed idx={idx}"));
}
let a_contig = Array::from_vec(a.to_vec()).reshape(&[m, k]);
assert!(a_contig.is_c_contiguous());
let c_b_only = a_contig.matmul(&b).expect("matmul should succeed");
for (idx, (got, want)) in c_b_only.to_vec().iter().zip(&expected).enumerate() {
assert_close_rel(*got, *want, 1e-9, format!("b transposed idx={idx}"));
}
}
#[test]
fn matmul_2d_handles_non_contiguous_operands_on_generic_tier() {
let a = Array::from_vec(vec![1i32, 2, 3, 4, 5, 6])
.reshape(&[2, 3])
.transpose_axis(0, 1);
assert_eq!(a.to_vec(), vec![1, 4, 2, 5, 3, 6]);
assert!(!a.is_c_contiguous());
let b = Array::from_vec(vec![1i32, 0, 0, 1]).reshape(&[2, 2]);
let c = a.matmul(&b).expect("matmul should succeed");
assert_eq!(c.shape(), vec![3, 2]);
assert_eq!(c.to_vec(), vec![1, 4, 2, 5, 3, 6]);
}
fn naive_batched_f64(batch: usize, m: usize, k: usize, n: usize, a: &[f64], b: &[f64]) -> Vec<f64> {
let mut out = Vec::with_capacity(batch * m * n);
for t in 0..batch {
out.extend(naive_gemm_f64(
m,
k,
n,
&a[t * m * k..(t + 1) * m * k],
&b[t * k * n..(t + 1) * k * n],
));
}
out
}
#[test]
fn matmul_3d_batched_matches_loop_of_2d_oracle() {
for &(batch, m, k, n) in &[
(1usize, 4usize, 5usize, 3usize),
(3, 8, 7, 6),
(8, 16, 16, 16),
(64, 16, 16, 16),
(5, 33, 65, 31),
] {
let a_data = seq_f64(batch * m * k, 0.0625, -1.25);
let b_data = seq_f64(batch * k * n, -0.03125, 2.5);
let expected = naive_batched_f64(batch, m, k, n, &a_data, &b_data);
let a = Array::from_vec(a_data).reshape(&[batch, m, k]);
let b = Array::from_vec(b_data).reshape(&[batch, k, n]);
let c = a.matmul(&b).expect("batched matmul should succeed");
assert_eq!(
c.shape(),
vec![batch, m, n],
"(batch={batch},m={m},k={k},n={n}) shape"
);
for (idx, (got, want)) in c.to_vec().iter().zip(&expected).enumerate() {
assert_close_rel(
*got,
*want,
1e-9,
format!("(batch={batch},m={m},k={k},n={n}) idx={idx}"),
);
}
}
}
#[test]
fn matmul_4d_batched_flattens_leading_axes_consistently() {
let (b0, b1, m, k, n) = (2usize, 3usize, 6usize, 5usize, 4usize);
let batch = b0 * b1;
let a_data = seq_f64(batch * m * k, 0.1, -0.5);
let b_data = seq_f64(batch * k * n, -0.2, 1.0);
let expected = naive_batched_f64(batch, m, k, n, &a_data, &b_data);
let a4 = Array::from_vec(a_data.clone()).reshape(&[b0, b1, m, k]);
let b4 = Array::from_vec(b_data.clone()).reshape(&[b0, b1, k, n]);
let c4 = a4.matmul(&b4).expect("4-D batched matmul should succeed");
assert_eq!(c4.shape(), vec![b0, b1, m, n]);
let a3 = Array::from_vec(a_data).reshape(&[batch, m, k]);
let b3 = Array::from_vec(b_data).reshape(&[batch, k, n]);
let c3 = a3.matmul(&b3).expect("3-D batched matmul should succeed");
assert_eq!(
c4.to_vec(),
c3.to_vec(),
"4-D and 3-D must agree bit-for-bit"
);
for (idx, (got, want)) in c4.to_vec().iter().zip(&expected).enumerate() {
assert_close_rel(*got, *want, 1e-9, format!("4-D idx={idx}"));
}
}
#[test]
fn matmul_3d_batched_generic_tier_matches_oracle() {
let (batch, m, k, n) = (4usize, 5usize, 3usize, 6usize);
let a_data = seq_i32(batch * m * k, 7, -3);
let b_data = seq_i32(batch * k * n, 5, -2);
let mut expected = Vec::with_capacity(batch * m * n);
for t in 0..batch {
expected.extend(naive_gemm_i32(
m,
k,
n,
&a_data[t * m * k..(t + 1) * m * k],
&b_data[t * k * n..(t + 1) * k * n],
));
}
let a = Array::from_vec(a_data).reshape(&[batch, m, k]);
let b = Array::from_vec(b_data).reshape(&[batch, k, n]);
let c = a.matmul(&b).expect("batched matmul should succeed");
assert_eq!(c.shape(), vec![batch, m, n]);
assert_eq!(c.to_vec(), expected);
}
#[test]
fn matmul_broadcast_batched_pairs_panels_correctly() {
let (m, k, n) = (4usize, 3usize, 5usize);
let a_data = seq_f64(2 * m * k, 0.25, -1.0);
let b_data = seq_f64(3 * k * n, -0.5, 2.0);
let a = Array::from_vec(a_data.clone()).reshape(&[2, 1, m, k]);
let b = Array::from_vec(b_data.clone()).reshape(&[1, 3, k, n]);
let c = a
.matmul(&b)
.expect("broadcast batched matmul should succeed");
assert_eq!(c.shape(), vec![2, 3, m, n]);
let got = c.to_vec();
for i in 0..2 {
for j in 0..3 {
let a_panel = &a_data[i * m * k..(i + 1) * m * k];
let b_panel = &b_data[j * k * n..(j + 1) * k * n];
let want = naive_gemm_f64(m, k, n, a_panel, b_panel);
let base = (i * 3 + j) * m * n;
for (idx, w) in want.iter().enumerate() {
assert_close_rel(got[base + idx], *w, 1e-9, format!("i={i} j={j} idx={idx}"));
}
}
}
}
#[test]
fn matmul_broadcast_batched_handles_rank_mismatch() {
let (m, k, n) = (3usize, 4usize, 2usize);
let a_data = seq_f64(3 * m * k, 0.5, -2.0);
let b_data = seq_f64(2 * 3 * k * n, -0.25, 1.0);
let a = Array::from_vec(a_data.clone()).reshape(&[3, m, k]);
let b = Array::from_vec(b_data.clone()).reshape(&[2, 3, k, n]);
let c = a.matmul(&b).expect("rank-broadcast matmul should succeed");
assert_eq!(c.shape(), vec![2, 3, m, n]);
let got = c.to_vec();
for i in 0..2 {
for j in 0..3 {
let a_panel = &a_data[j * m * k..(j + 1) * m * k];
let b_panel = &b_data[(i * 3 + j) * k * n..(i * 3 + j + 1) * k * n];
let want = naive_gemm_f64(m, k, n, a_panel, b_panel);
let base = (i * 3 + j) * m * n;
for (idx, w) in want.iter().enumerate() {
assert_close_rel(got[base + idx], *w, 1e-9, format!("i={i} j={j} idx={idx}"));
}
}
}
}
#[test]
fn matmul_respects_identity_and_associativity() {
let n = 65usize;
let mut eye = vec![0.0f64; n * n];
for i in 0..n {
eye[i * n + i] = 1.0;
}
let identity = Array::from_vec(eye).reshape(&[n, n]);
let a = Array::from_vec(seq_f64(n * n, 0.01, -0.5)).reshape(&[n, n]);
let b = Array::from_vec(seq_f64(n * n, -0.02, 0.25)).reshape(&[n, n]);
let ai = a.matmul(&identity).expect("matmul should succeed");
for (idx, (got, want)) in ai.to_vec().iter().zip(a.to_vec()).enumerate() {
assert_close_rel(*got, want, 1e-12, format!("A*I idx={idx}"));
}
let ia = identity.matmul(&a).expect("matmul should succeed");
for (idx, (got, want)) in ia.to_vec().iter().zip(a.to_vec()).enumerate() {
assert_close_rel(*got, want, 1e-12, format!("I*A idx={idx}"));
}
let left = a
.matmul(&b)
.expect("matmul should succeed")
.matmul(&identity)
.expect("matmul should succeed");
let right = a
.matmul(&b.matmul(&identity).expect("matmul should succeed"))
.expect("matmul should succeed");
for (idx, (got, want)) in left.to_vec().iter().zip(right.to_vec()).enumerate() {
assert_close_rel(*got, want, 1e-9, format!("assoc idx={idx}"));
}
}
#[test]
fn matmul_rejects_incompatible_inner_dimensions() {
let a = Array::from_vec(seq_f64(6, 1.0, 0.0)).reshape(&[2, 3]);
let b = Array::from_vec(seq_f64(8, 1.0, 0.0)).reshape(&[4, 2]);
assert!(
a.matmul(&b).is_err(),
"3 != 4 inner dimensions must be rejected"
);
let a3 = Array::from_vec(seq_f64(12, 1.0, 0.0)).reshape(&[2, 2, 3]);
let b3 = Array::from_vec(seq_f64(16, 1.0, 0.0)).reshape(&[2, 4, 2]);
assert!(
a3.matmul(&b3).is_err(),
"batched 3 != 4 inner dimensions must be rejected"
);
}
mod perf {
use super::*;
use std::time::{Duration, Instant};
pub fn legacy_matmul_2d_blocked(a: &Array<f64>, b: &Array<f64>) -> Array<f64> {
let a_shape = a.shape();
let b_shape = b.shape();
let m = a_shape[0];
let k = a_shape[1];
let n = b_shape[1];
let mut c_data = vec![0.0f64; m * n];
let owned_a;
let a_data: &[f64] = match a.as_slice() {
Some(slice) => slice,
None => {
owned_a = a.to_vec();
&owned_a
}
};
let owned_b;
let b_data: &[f64] = match b.as_slice() {
Some(slice) => slice,
None => {
owned_b = b.to_vec();
&owned_b
}
};
const BLOCK_SIZE: usize = 64;
for i_block in (0..m).step_by(BLOCK_SIZE) {
for k_block in (0..k).step_by(BLOCK_SIZE) {
for j_block in (0..n).step_by(BLOCK_SIZE) {
let i_end = std::cmp::min(i_block + BLOCK_SIZE, m);
let k_end = std::cmp::min(k_block + BLOCK_SIZE, k);
let j_end = std::cmp::min(j_block + BLOCK_SIZE, n);
for i in i_block..i_end {
for k_l in k_block..k_end {
let a_ik = a_data[i * k + k_l];
for j in j_block..j_end {
c_data[i * n + j] += a_ik * b_data[k_l * n + j];
}
}
}
}
}
}
Array::from_vec(c_data).reshape(&[m, n])
}
fn alternating_min<A, B, RA, RB>(reps: usize, mut a: A, mut b: B) -> (Duration, Duration)
where
A: FnMut() -> RA,
B: FnMut() -> RB,
{
std::hint::black_box(a());
std::hint::black_box(b());
let mut min_a = Duration::MAX;
let mut min_b = Duration::MAX;
for _ in 0..reps {
let t = Instant::now();
std::hint::black_box(a());
min_a = min_a.min(t.elapsed());
let t = Instant::now();
std::hint::black_box(b());
min_b = min_b.min(t.elapsed());
}
(min_a, min_b)
}
fn mat(m: usize, n: usize) -> Array<f64> {
Array::from_vec(
(0..m * n)
.map(|i| (i as f64) * 0.125 - 3.0)
.collect::<Vec<_>>(),
)
.reshape(&[m, n])
}
#[test]
#[ignore = "performance measurement, not a correctness assertion"]
fn matmul_perf_evidence_2d() {
println!(
"\n{:<16} {:>14} {:>14} {:>10}",
"shape (m,k,n)", "dispatched", "legacy", "speedup"
);
for &(m, k, n, reps) in &[
(8usize, 8usize, 8usize, 2000usize),
(32, 32, 32, 500),
(64, 64, 64, 200),
(128, 128, 128, 100),
(256, 256, 256, 40),
(512, 512, 512, 15),
(512, 64, 512, 30),
] {
let a = mat(m, k);
let b = mat(k, n);
let (d, l) = alternating_min(
reps,
|| a.matmul(&b).expect("matmul should succeed"),
|| legacy_matmul_2d_blocked(&a, &b),
);
println!(
"{:<16} {:>14?} {:>14?} {:>9.2}x",
format!("{m},{k},{n}"),
d,
l,
l.as_secs_f64() / d.as_secs_f64()
);
}
}
#[test]
#[ignore = "performance measurement, not a correctness assertion"]
fn matmul_perf_evidence_batched() {
println!(
"\n{:<20} {:>14} {:>14} {:>10}",
"batch/panel", "dispatched", "legacy_ixdyn", "speedup"
);
for &(batch, panel, reps) in &[
(1usize, 16usize, 300usize),
(8, 16, 200),
(64, 16, 60),
(1, 64, 60),
(8, 64, 30),
(64, 64, 8),
] {
let a = Array::from_vec(
(0..batch * panel * panel)
.map(|i| (i as f64) * 0.125 - 3.0)
.collect::<Vec<_>>(),
)
.reshape(&[batch, panel, panel]);
let b = a.clone();
let (d, l) = alternating_min(
reps,
|| a.matmul(&b).expect("matmul should succeed"),
|| legacy_batched_ixdyn(&a, &b),
);
println!(
"{:<20} {:>14?} {:>14?} {:>9.2}x",
format!("batch{batch}/panel{panel}"),
d,
l,
l.as_secs_f64() / d.as_secs_f64()
);
}
}
fn legacy_batched_ixdyn(a: &Array<f64>, b: &Array<f64>) -> Array<f64> {
use scirs2_core::ndarray::IxDyn;
let a_shape = a.shape();
let b_shape = b.shape();
let batch_shape = &a_shape[..a_shape.len() - 2];
let m = a_shape[a_shape.len() - 2];
let k = a_shape[a_shape.len() - 1];
let n = b_shape[b_shape.len() - 1];
let mut output_shape = batch_shape.to_vec();
output_shape.push(m);
output_shape.push(n);
let mut result = Array::<f64>::zeros(&output_shape);
let batch_size: usize = batch_shape.iter().product();
for batch_idx in 0..batch_size {
let mut batch_indices = Vec::with_capacity(batch_shape.len());
let mut temp = batch_idx;
for &dim in batch_shape.iter().rev() {
batch_indices.insert(0, temp % dim);
temp /= dim;
}
let mut a_indices = batch_indices.clone();
a_indices.push(0);
a_indices.push(0);
let mut b_indices = batch_indices.clone();
b_indices.push(0);
b_indices.push(0);
for i in 0..m {
let p = a_indices.len() - 2;
a_indices[p] = i;
for j in 0..n {
let p = b_indices.len() - 1;
b_indices[p] = j;
let mut sum = 0.0f64;
for l in 0..k {
let p = a_indices.len() - 1;
a_indices[p] = l;
let p = b_indices.len() - 2;
b_indices[p] = l;
sum += a.array().get(IxDyn(&a_indices)).expect("valid index")
* b.array().get(IxDyn(&b_indices)).expect("valid index");
}
let mut out = batch_indices.clone();
out.push(i);
out.push(j);
result.set(&out, sum).expect("valid output index");
}
}
}
result
}
}
mod bakeoff {
use scirs2_core::ndarray::linalg::general_mat_mul;
use scirs2_core::ndarray::{ArrayView2, ArrayViewMut2};
use scirs2_core::parallel_ops::*;
use std::time::{Duration, Instant};
pub fn blocked_serial_f64(m: usize, k: usize, n: usize, a: &[f64], b: &[f64], c: &mut [f64]) {
for v in c.iter_mut() {
*v = 0.0;
}
const BLOCK_SIZE: usize = 64;
for i_block in (0..m).step_by(BLOCK_SIZE) {
for k_block in (0..k).step_by(BLOCK_SIZE) {
for j_block in (0..n).step_by(BLOCK_SIZE) {
let i_end = std::cmp::min(i_block + BLOCK_SIZE, m);
let k_end = std::cmp::min(k_block + BLOCK_SIZE, k);
let j_end = std::cmp::min(j_block + BLOCK_SIZE, n);
for i in i_block..i_end {
for k_l in k_block..k_end {
let a_ik = a[i * k + k_l];
for j in j_block..j_end {
c[i * n + j] += a_ik * b[k_l * n + j];
}
}
}
}
}
}
}
pub fn blocked_serial_f32(m: usize, k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
for v in c.iter_mut() {
*v = 0.0;
}
const BLOCK_SIZE: usize = 64;
for i_block in (0..m).step_by(BLOCK_SIZE) {
for k_block in (0..k).step_by(BLOCK_SIZE) {
for j_block in (0..n).step_by(BLOCK_SIZE) {
let i_end = std::cmp::min(i_block + BLOCK_SIZE, m);
let k_end = std::cmp::min(k_block + BLOCK_SIZE, k);
let j_end = std::cmp::min(j_block + BLOCK_SIZE, n);
for i in i_block..i_end {
for k_l in k_block..k_end {
let a_ik = a[i * k + k_l];
for j in j_block..j_end {
c[i * n + j] += a_ik * b[k_l * n + j];
}
}
}
}
}
}
}
pub fn blocked_par_f64(m: usize, k: usize, n: usize, a: &[f64], b: &[f64], c: &mut [f64]) {
if m == 0 || k == 0 || n == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let chunk_rows = m.div_ceil(current_num_threads().max(1)).max(1);
a.par_chunks(chunk_rows * k)
.zip(c.par_chunks_mut(chunk_rows * n))
.for_each(|(a_chunk, c_chunk)| {
let rows = a_chunk.len() / k;
if rows == 0 {
return;
}
blocked_serial_f64(rows, k, n, a_chunk, b, c_chunk);
});
}
pub fn blocked_par_f32(m: usize, k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
if m == 0 || k == 0 || n == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let chunk_rows = m.div_ceil(current_num_threads().max(1)).max(1);
a.par_chunks(chunk_rows * k)
.zip(c.par_chunks_mut(chunk_rows * n))
.for_each(|(a_chunk, c_chunk)| {
let rows = a_chunk.len() / k;
if rows == 0 {
return;
}
blocked_serial_f32(rows, k, n, a_chunk, b, c_chunk);
});
}
pub fn simd_serial_f64(m: usize, k: usize, n: usize, a: &[f64], b: &[f64], c: &mut [f64]) {
scirs2_core::simd_ops::simd_matrix_multiply_f64(m, k, n, 1.0, a, b, 0.0, c);
}
pub fn simd_serial_f32(m: usize, k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
scirs2_core::simd_ops::simd_matrix_multiply_f32(m, k, n, 1.0, a, b, 0.0, c);
}
pub fn simd_par_f64(m: usize, k: usize, n: usize, a: &[f64], b: &[f64], c: &mut [f64]) {
if m == 0 || k == 0 || n == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let chunk_rows = m.div_ceil(current_num_threads().max(1)).max(1);
a.par_chunks(chunk_rows * k)
.zip(c.par_chunks_mut(chunk_rows * n))
.for_each(|(a_chunk, c_chunk)| {
let rows = a_chunk.len() / k;
if rows == 0 {
return;
}
scirs2_core::simd_ops::simd_matrix_multiply_f64(
rows, k, n, 1.0, a_chunk, b, 0.0, c_chunk,
);
});
}
pub fn simd_par_f32(m: usize, k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
if m == 0 || k == 0 || n == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let chunk_rows = m.div_ceil(current_num_threads().max(1)).max(1);
a.par_chunks(chunk_rows * k)
.zip(c.par_chunks_mut(chunk_rows * n))
.for_each(|(a_chunk, c_chunk)| {
let rows = a_chunk.len() / k;
if rows == 0 {
return;
}
scirs2_core::simd_ops::simd_matrix_multiply_f32(
rows, k, n, 1.0, a_chunk, b, 0.0, c_chunk,
);
});
}
pub fn blas_acc_f64(m: usize, k: usize, n: usize, a: &[f64], b: &[f64], c: &mut [f64]) {
let a_view = ArrayView2::from_shape((m, k), a).expect("operand shape should match slice");
let b_view = ArrayView2::from_shape((k, n), b).expect("operand shape should match slice");
let out = scirs2_linalg::blas_accelerated::matmul(&a_view, &b_view)
.expect("blas_accelerated matmul should succeed");
for (dst, src) in c.iter_mut().zip(out.iter()) {
*dst = *src;
}
}
pub fn blas_acc_f32(m: usize, k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
let a_view = ArrayView2::from_shape((m, k), a).expect("operand shape should match slice");
let b_view = ArrayView2::from_shape((k, n), b).expect("operand shape should match slice");
let out = scirs2_linalg::blas_accelerated::matmul(&a_view, &b_view)
.expect("blas_accelerated matmul should succeed");
for (dst, src) in c.iter_mut().zip(out.iter()) {
*dst = *src;
}
}
pub fn nd_gmm_f64(m: usize, k: usize, n: usize, a: &[f64], b: &[f64], c: &mut [f64]) {
if k == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let a_view = ArrayView2::from_shape((m, k), a).expect("operand shape should match slice");
let b_view = ArrayView2::from_shape((k, n), b).expect("operand shape should match slice");
let mut c_view =
ArrayViewMut2::from_shape((m, n), c).expect("output shape should match slice");
general_mat_mul(1.0, &a_view, &b_view, 0.0, &mut c_view);
}
pub fn nd_gmm_f32(m: usize, k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
if k == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let a_view = ArrayView2::from_shape((m, k), a).expect("operand shape should match slice");
let b_view = ArrayView2::from_shape((k, n), b).expect("operand shape should match slice");
let mut c_view =
ArrayViewMut2::from_shape((m, n), c).expect("output shape should match slice");
general_mat_mul(1.0, &a_view, &b_view, 0.0, &mut c_view);
}
pub fn nd_gmm_par_f64(m: usize, k: usize, n: usize, a: &[f64], b: &[f64], c: &mut [f64]) {
if m == 0 || k == 0 || n == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let chunk_rows = m.div_ceil(current_num_threads().max(1)).max(1);
a.par_chunks(chunk_rows * k)
.zip(c.par_chunks_mut(chunk_rows * n))
.for_each(|(a_chunk, c_chunk)| {
let rows = a_chunk.len() / k;
if rows == 0 {
return;
}
nd_gmm_f64(rows, k, n, a_chunk, b, c_chunk);
});
}
pub fn nd_gmm_par_f32(m: usize, k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
if m == 0 || k == 0 || n == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let chunk_rows = m.div_ceil(current_num_threads().max(1)).max(1);
a.par_chunks(chunk_rows * k)
.zip(c.par_chunks_mut(chunk_rows * n))
.for_each(|(a_chunk, c_chunk)| {
let rows = a_chunk.len() / k;
if rows == 0 {
return;
}
nd_gmm_f32(rows, k, n, a_chunk, b, c_chunk);
});
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Backend {
BlockedSerial,
BlockedPar,
SimdSerial,
SimdPar,
BlasAcc,
NdGmm,
NdGmmPar,
}
const BACKENDS: &[Backend] = &[
Backend::BlockedSerial,
Backend::BlockedPar,
Backend::SimdSerial,
Backend::SimdPar,
Backend::BlasAcc,
Backend::NdGmm,
Backend::NdGmmPar,
];
const BACKEND_NAMES: &[&str] = &[
"blk_ser",
"blk_par",
"simd_ser",
"simd_par",
"blas_acc",
"nd_gmm",
"nd_gmm_par",
];
fn run_f64(be: Backend, m: usize, k: usize, n: usize, a: &[f64], b: &[f64], c: &mut [f64]) {
match be {
Backend::BlockedSerial => blocked_serial_f64(m, k, n, a, b, c),
Backend::BlockedPar => blocked_par_f64(m, k, n, a, b, c),
Backend::SimdSerial => simd_serial_f64(m, k, n, a, b, c),
Backend::SimdPar => simd_par_f64(m, k, n, a, b, c),
Backend::BlasAcc => blas_acc_f64(m, k, n, a, b, c),
Backend::NdGmm => nd_gmm_f64(m, k, n, a, b, c),
Backend::NdGmmPar => nd_gmm_par_f64(m, k, n, a, b, c),
}
}
fn run_f32(be: Backend, m: usize, k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
match be {
Backend::BlockedSerial => blocked_serial_f32(m, k, n, a, b, c),
Backend::BlockedPar => blocked_par_f32(m, k, n, a, b, c),
Backend::SimdSerial => simd_serial_f32(m, k, n, a, b, c),
Backend::SimdPar => simd_par_f32(m, k, n, a, b, c),
Backend::BlasAcc => blas_acc_f32(m, k, n, a, b, c),
Backend::NdGmm => nd_gmm_f32(m, k, n, a, b, c),
Backend::NdGmmPar => nd_gmm_par_f32(m, k, n, a, b, c),
}
}
fn reps_for(flops: f64) -> usize {
let want = 6.0e8 / flops.max(1.0);
want.clamp(7.0, 1500.0) as usize
}
fn seq_f64(len: usize, step: f64, base: f64) -> Vec<f64> {
(0..len).map(|i| (i as f64) * step + base).collect()
}
fn seq_f32(len: usize, step: f32, base: f32) -> Vec<f32> {
(0..len).map(|i| (i as f32) * step + base).collect()
}
fn measure_f64(m: usize, k: usize, n: usize) -> Vec<Duration> {
let a = seq_f64(m * k, 0.0125, -3.0);
let b = seq_f64(k * n, -0.00625, 1.5);
let mut c = vec![0.0f64; m * n];
for &be in BACKENDS {
run_f64(be, m, k, n, &a, &b, &mut c);
std::hint::black_box(&c);
}
let flops = 2.0 * (m as f64) * (k as f64) * (n as f64);
let reps = reps_for(flops);
let mut best = vec![Duration::MAX; BACKENDS.len()];
for rep in 0..reps {
for offset in 0..BACKENDS.len() {
let idx = (offset + rep) % BACKENDS.len();
let t = Instant::now();
run_f64(BACKENDS[idx], m, k, n, &a, &b, &mut c);
let dt = t.elapsed();
std::hint::black_box(&c);
best[idx] = best[idx].min(dt);
}
}
best
}
fn measure_f32(m: usize, k: usize, n: usize) -> Vec<Duration> {
let a = seq_f32(m * k, 0.0125, -3.0);
let b = seq_f32(k * n, -0.00625, 1.5);
let mut c = vec![0.0f32; m * n];
for &be in BACKENDS {
run_f32(be, m, k, n, &a, &b, &mut c);
std::hint::black_box(&c);
}
let flops = 2.0 * (m as f64) * (k as f64) * (n as f64);
let reps = reps_for(flops);
let mut best = vec![Duration::MAX; BACKENDS.len()];
for rep in 0..reps {
for offset in 0..BACKENDS.len() {
let idx = (offset + rep) % BACKENDS.len();
let t = Instant::now();
run_f32(BACKENDS[idx], m, k, n, &a, &b, &mut c);
let dt = t.elapsed();
std::hint::black_box(&c);
best[idx] = best[idx].min(dt);
}
}
best
}
fn print_header(tag: &str) {
print!("{tag:<10} {:>16}", "m,k,n");
for name in BACKEND_NAMES {
print!(" {name:>12}");
}
println!(" {:>12}", "winner");
}
fn print_row(tag: &str, m: usize, k: usize, n: usize, best: &[Duration]) {
print!("{tag:<10} {:>16}", format!("{m},{k},{n}"));
for d in best {
print!(" {:>12}", d.as_nanos());
}
let mut win = 0usize;
for (i, d) in best.iter().enumerate() {
if *d < best[win] {
win = i;
}
}
println!(" {:>12}", BACKEND_NAMES[win]);
}
fn naive(m: usize, k: usize, n: usize, a: &[f64], b: &[f64]) -> Vec<f64> {
let mut c = vec![0.0f64; m * n];
for i in 0..m {
for p in 0..k {
let a_ip = a[i * k + p];
for j in 0..n {
c[i * n + j] += a_ip * b[p * n + j];
}
}
}
c
}
#[test]
fn every_bakeoff_candidate_matches_naive_oracle() {
for &(m, k, n) in &[
(1usize, 1usize, 1usize),
(5, 0, 4), (0, 5, 4), (5, 4, 0), (3, 7, 5), (33, 17, 41),
(64, 64, 64),
(65, 32, 129),
(128, 16, 96),
] {
let a = seq_f64(m * k, 0.0125, -3.0);
let b = seq_f64(k * n, -0.00625, 1.5);
let expected = naive(m, k, n, &a, &b);
for (idx, &be) in BACKENDS.iter().enumerate() {
let mut c = vec![f64::NAN; m * n];
run_f64(be, m, k, n, &a, &b, &mut c);
for (i, (&got, &want)) in c.iter().zip(expected.iter()).enumerate() {
let scale = want.abs().max(1.0);
assert!(
(got - want).abs() <= 1e-9 * scale,
"{} at (m={m},k={k},n={n}) idx {i}: got {got}, want {want}",
BACKEND_NAMES[idx]
);
}
}
}
}
#[test]
#[ignore = "performance measurement, not a correctness assertion"]
fn bakeoff_square_f64() {
print_header("SQ64");
for &s in &[
8usize, 16, 32, 48, 64, 80, 96, 112, 128, 160, 192, 256, 320, 384, 512,
] {
let best = measure_f64(s, s, s);
print_row("SQ64", s, s, s, &best);
}
}
#[test]
#[ignore = "performance measurement, not a correctness assertion"]
fn bakeoff_rectangular_f64() {
print_header("RECT64");
for &k in &[8usize, 16, 24, 32, 48, 64, 96, 128, 192, 256, 512] {
let best = measure_f64(512, k, 512);
print_row("RECT64", 512, k, 512, &best);
}
for &(m, k, n) in &[
(64usize, 512usize, 64usize),
(256, 32, 256),
(1024, 32, 1024),
] {
let best = measure_f64(m, k, n);
print_row("RECT64", m, k, n, &best);
}
}
#[test]
#[ignore = "performance measurement, not a correctness assertion"]
fn bakeoff_f32() {
print_header("F32");
for &s in &[32usize, 48, 64, 96, 128, 192, 256, 384] {
let best = measure_f32(s, s, s);
print_row("F32", s, s, s, &best);
}
for &(m, k, n) in &[(512usize, 64usize, 512usize), (512, 32, 512)] {
let best = measure_f32(m, k, n);
print_row("F32", m, k, n, &best);
}
}
}