#[cfg(feature = "machine_learning")]
use ndarray::{Array1, Ix1};
#[cfg(any(feature = "machine_learning", feature = "neural_network"))]
use ndarray::{Array2, ArrayBase, Data, Ix2};
#[cfg(any(feature = "machine_learning", feature = "neural_network"))]
pub(crate) fn dot_par<T, S1, S2>(
a: &ArrayBase<S1, Ix2>,
b: &ArrayBase<S2, Ix2>,
par: gemmkit_ndarray::Parallelism,
) -> Array2<T>
where
T: gemmkit_ndarray::GemmScalar,
S1: Data<Elem = T>,
S2: Data<Elem = T>,
{
let (m, k) = a.dim();
let (kb, n) = b.dim();
assert_eq!(
k, kb,
"dot_par: inner dimensions disagree (a is {m}x{k}, b is {kb}x{n})",
);
let mut c = Array2::from_elem((m, n), T::ZERO);
gemmkit_ndarray::gemm(T::ONE, a, b, T::ZERO, &mut c, par);
c
}
#[cfg(feature = "machine_learning")]
pub(crate) fn matvec<T, S1, S2>(
a: &ArrayBase<S1, Ix2>,
x: &ArrayBase<S2, Ix1>,
par: gemmkit_ndarray::Parallelism,
) -> Array1<T>
where
T: gemmkit_ndarray::GemmScalar,
S1: Data<Elem = T>,
S2: Data<Elem = T>,
{
use ndarray::Axis;
let (m, k) = a.dim();
assert_eq!(
k,
x.len(),
"matvec: inner dimensions disagree (a is {m}x{k}, x has length {})",
x.len()
);
let x_col = x.view().insert_axis(Axis(1)); let mut y = Array1::from_elem(m, T::ZERO);
let mut y_col = y.view_mut().insert_axis(Axis(1)); gemmkit_ndarray::gemm(T::ONE, a, &x_col, T::ZERO, &mut y_col, par);
y
}
tunable_gate! {
pub(crate) GEMM_CHUNK_ELEMS => gemm_chunk_elems / set_gemm_chunk_elems = 33_554_432
}
#[doc(hidden)]
pub fn gemm_chunk_rows(row_len: usize) -> usize {
(gemm_chunk_elems() / row_len.max(1)).clamp(16, 4096)
}
tunable_gate! {
pub(crate) CACHE_RESIDENT_MAX_BYTES
=> cache_resident_max_bytes / set_cache_resident_max_bytes = 64 * 1024 * 1024
}
#[doc(hidden)]
pub fn cache_resident<T>(rows: usize, cols: usize) -> bool {
rows.saturating_mul(cols)
.saturating_mul(std::mem::size_of::<T>())
< cache_resident_max_bytes()
}
#[cfg(all(test, any(feature = "machine_learning", feature = "neural_network")))]
mod tests {
use super::*;
use gemmkit_ndarray::Parallelism;
use gemmkit_ndarray::dot;
#[cfg(feature = "machine_learning")]
use ndarray::Array1;
use ndarray::{Array2, LinalgScalar, s};
fn rand_f32(r: usize, c: usize, seed: u64) -> Array2<f32> {
Array2::from_shape_fn((r, c), |(i, j)| {
let t = (seed as f64) * 0.731 + (i * c + j) as f64 * 0.618_033_988_7;
((t.sin() * 43758.5453).fract() - 0.5) as f32
})
}
fn rand_f64(r: usize, c: usize, seed: u64) -> Array2<f64> {
Array2::from_shape_fn((r, c), |(i, j)| {
let t = (seed as f64) * 0.731 + (i * c + j) as f64 * 0.618_033_988_7;
(t.sin() * 43758.5453).fract() - 0.5
})
}
fn naive<T: LinalgScalar>(a: &Array2<T>, b: &Array2<T>) -> Array2<T> {
let (m, k) = a.dim();
let n = b.ncols();
let mut c = Array2::<T>::zeros((m, n));
for i in 0..m {
for j in 0..n {
let mut acc = T::zero();
for p in 0..k {
acc = acc + a[[i, p]] * b[[p, j]];
}
c[[i, j]] = acc;
}
}
c
}
fn assert_close_f32(got: &Array2<f32>, want: &Array2<f32>, eps: f32) {
assert_eq!(got.shape(), want.shape());
for (g, w) in got.iter().zip(want.iter()) {
assert!((g - w).abs() <= eps, "f32 mismatch: {g} vs {w}");
}
}
fn assert_close_f64(got: &Array2<f64>, want: &Array2<f64>, eps: f64) {
assert_eq!(got.shape(), want.shape());
for (g, w) in got.iter().zip(want.iter()) {
assert!((g - w).abs() <= eps, "f64 mismatch: {g} vs {w}");
}
}
#[test]
fn dot_matches_reference_f32() {
for &(m, k, n) in &[
(17usize, 23usize, 19usize),
(64, 48, 64),
] {
let a = rand_f32(m, k, 1);
let b = rand_f32(k, n, 2);
assert_close_f32(&dot(&a, &b), &naive(&a, &b), 1e-2);
}
let a = rand_f32(256, 300, 3);
let b = rand_f32(300, 256, 4);
assert_close_f32(&dot(&a, &b), &a.dot(&b), 1e-2);
}
#[test]
fn dot_matches_reference_f64() {
for &(m, k, n) in &[
(17usize, 23usize, 19usize),
(64, 64, 64),
] {
let a = rand_f64(m, k, 1);
let b = rand_f64(k, n, 2);
assert_close_f64(&dot(&a, &b), &naive(&a, &b), 1e-9);
}
let a = rand_f64(128, 128, 3);
let b = rand_f64(128, 128, 4);
assert_close_f64(&dot(&a, &b), &a.dot(&b), 1e-9);
}
#[test]
fn dot_strided_operands() {
let a = rand_f64(40, 24, 5);
let b = rand_f64(40, 18, 6);
let got = dot(&a.t(), &b);
let want = naive(&a.t().to_owned(), &b);
assert_close_f64(&got, &want, 1e-9);
let a = rand_f64(40, 30, 7);
let b = rand_f64(30, 40, 8);
let a_sl = a.slice(s![..;2, ..]); let b_sl = b.slice(s![.., ..;2]); let got = dot(&a_sl, &b_sl);
let want = naive(&a_sl.to_owned(), &b_sl.to_owned());
assert_close_f64(&got, &want, 1e-9);
}
#[test]
fn dot_thin_output() {
let a = rand_f64(4096, 64, 61);
let b = rand_f64(64, 4, 62);
assert_close_f64(&dot(&a, &b), &a.dot(&b), 1e-9);
let a = rand_f32(16384, 64, 63);
let b = rand_f32(64, 4, 64);
assert_close_f32(&dot(&a, &b), &a.dot(&b), 1e-2);
}
#[test]
fn dot_par_thread_count_independent_f64() {
for &(m, k, n) in &[(96usize, 96usize, 96usize), (256, 64, 64), (64, 8192, 64)] {
let a = rand_f64(m, k, 11);
let b = rand_f64(k, n, 12);
let serial = dot_par(&a, &b, Parallelism::Serial);
for threads in [2usize, 4, 8, 16, 32] {
let par = dot_par(&a, &b, Parallelism::Rayon(threads));
assert!(
serial
.iter()
.zip(par.iter())
.all(|(s, p)| s.to_bits() == p.to_bits()),
"gemm f64 {m}x{k}x{n} differs between serial and Rayon({threads})"
);
}
}
}
#[test]
fn dot_par_thread_count_independent_f32() {
for &(m, k, n) in &[(96usize, 96usize, 96usize), (64, 8192, 64)] {
let a = rand_f32(m, k, 13);
let b = rand_f32(k, n, 14);
let serial = dot_par(&a, &b, Parallelism::Serial);
for threads in [2usize, 4, 32] {
let par = dot_par(&a, &b, Parallelism::Rayon(threads));
assert!(
serial
.iter()
.zip(par.iter())
.all(|(s, p)| s.to_bits() == p.to_bits()),
"gemm f32 {m}x{k}x{n} differs between serial and Rayon({threads})"
);
}
}
}
#[test]
fn dot_run_to_run_deterministic() {
let a = rand_f64(200, 200, 21);
let b = rand_f64(200, 200, 22);
let c1 = dot(&a, &b);
let c2 = dot(&a, &b);
assert!(
c1.iter()
.zip(c2.iter())
.all(|(x, y)| x.to_bits() == y.to_bits())
);
}
#[test]
fn gemm_fused_bias_relu_bitwise_matches_unfused() {
let a = rand_f32(64, 48, 31);
let b = rand_f32(48, 40, 32);
let bias: Vec<f32> = (0..40).map(|j| (j as f32) * 0.05 - 1.0).collect();
let mut fused = Array2::from_elem((64, 40), 0.0f32);
gemmkit_ndarray::gemm_fused(
1.0,
&a,
&b,
0.0,
&mut fused,
Some(gemmkit_ndarray::Bias::PerCol(&bias)),
Some(gemmkit_ndarray::Activation::Relu),
Parallelism::Rayon(0),
);
let mut want = dot(&a, &b);
for mut row in want.rows_mut() {
for (v, bj) in row.iter_mut().zip(bias.iter()) {
*v = (*v + bj).max(0.0);
}
}
assert!(
fused
.iter()
.zip(want.iter())
.all(|(f, w)| f.to_bits() == w.to_bits()),
"fused bias+ReLU differs from the unfused pipeline"
);
}
#[cfg(feature = "machine_learning")]
#[test]
fn matvec_matches_reference() {
for &(m, k) in &[(40usize, 24usize), (8192, 64)] {
let a = rand_f64(m, k, 31);
let x = Array1::from_shape_fn(k, |i| ((i as f64) * 0.37).sin());
let got = matvec(&a, &x, Parallelism::Rayon(0));
let want = a.dot(&x);
assert_eq!(got.len(), want.len());
for (g, w) in got.iter().zip(want.iter()) {
assert!((g - w).abs() <= 1e-9, "matvec mismatch: {g} vs {w}");
}
}
}
#[cfg(feature = "machine_learning")]
#[test]
fn matvec_serial_and_auto_agree_bitwise() {
let a = rand_f64(8192, 64, 41);
let x = Array1::from_shape_fn(64, |i| ((i as f64) * 0.59).sin());
let serial = matvec(&a, &x, Parallelism::Serial);
let auto = matvec(&a, &x, Parallelism::Rayon(0));
assert_eq!(serial.len(), auto.len());
assert!(
serial
.iter()
.zip(auto.iter())
.all(|(s, p)| s.to_bits() == p.to_bits()),
"matvec serial vs auto differ"
);
}
#[cfg(feature = "machine_learning")]
#[test]
fn matvec_run_to_run_deterministic() {
let a = rand_f64(8192, 64, 51);
let x = Array1::from_shape_fn(64, |i| ((i as f64) * 0.23).sin());
let y1 = matvec(&a, &x, Parallelism::Rayon(0));
let y2 = matvec(&a, &x, Parallelism::Rayon(0));
assert!(
y1.iter()
.zip(y2.iter())
.all(|(a, b)| a.to_bits() == b.to_bits())
);
}
#[test]
fn dot_edge_cases() {
let a = Array2::<f64>::zeros((0, 4));
let b = Array2::<f64>::zeros((4, 3));
assert_eq!(dot(&a, &b).shape(), &[0, 3]);
let a = Array2::<f64>::zeros((3, 0));
let b = Array2::<f64>::zeros((0, 4));
let c = dot(&a, &b);
assert_eq!(c.shape(), &[3, 4]);
assert!(c.iter().all(|&x| x == 0.0));
let a = Array2::<f64>::from_elem((1, 1), 3.0);
let b = Array2::<f64>::from_elem((1, 1), 4.0);
assert_eq!(dot(&a, &b)[[0, 0]], 12.0);
}
#[test]
#[should_panic]
fn dot_par_dimension_mismatch_panics() {
let a = Array2::<f64>::zeros((2, 3));
let b = Array2::<f64>::zeros((4, 2));
let _ = dot_par(&a, &b, Parallelism::Serial);
}
}