use crate::GemmTask;
const FLOP_THRESHOLD: usize = 1 << 13;
pub(super) fn gemm_f32(task: &GemmTask<'_, f32>) -> Option<Vec<f32>> {
if flops(task.m(), task.n(), task.k()) < FLOP_THRESHOLD {
return None;
}
executed_f32(task)
}
pub(super) fn gemm_f64(task: &GemmTask<'_, f64>) -> Option<Vec<f64>> {
if flops(task.m(), task.n(), task.k()) < FLOP_THRESHOLD {
return None;
}
executed_f64(task)
}
fn flops(m: usize, n: usize, k: usize) -> usize {
2usize.saturating_mul(m).saturating_mul(n).saturating_mul(k)
}
fn strides_for(strides: [usize; 2], rows: usize, columns: usize) -> Option<[isize; 2]> {
if strides[0] == 0 || strides[1] == 0 {
return None;
}
let row_stride = isize::try_from(strides[0]).ok()?;
let column_stride = isize::try_from(strides[1]).ok()?;
let furthest = (rows - 1)
.checked_mul(strides[0])?
.checked_add((columns - 1).checked_mul(strides[1])?)?;
isize::try_from(furthest).ok()?;
Some([row_stride, column_stride])
}
pub(super) fn executed_f32(task: &GemmTask<'_, f32>) -> Option<Vec<f32>> {
let a = strides_for(task.a_strides(), task.m(), task.k())?;
let b = strides_for(task.b_strides(), task.k(), task.n())?;
let volume = task.m().checked_mul(task.n())?;
isize::try_from(volume).ok()?;
let mut product = vec![0.0_f32; volume];
unsafe {
matrixmultiply::sgemm(
task.m(),
task.k(),
task.n(),
1.0,
task.a().as_ptr(),
a[0],
a[1],
task.b().as_ptr(),
b[0],
b[1],
0.0,
product.as_mut_ptr(),
task.n() as isize,
1,
);
}
Some(product)
}
pub(super) fn executed_f64(task: &GemmTask<'_, f64>) -> Option<Vec<f64>> {
let a = strides_for(task.a_strides(), task.m(), task.k())?;
let b = strides_for(task.b_strides(), task.k(), task.n())?;
let volume = task.m().checked_mul(task.n())?;
isize::try_from(volume).ok()?;
let mut product = vec![0.0_f64; volume];
unsafe {
matrixmultiply::dgemm(
task.m(),
task.k(),
task.n(),
1.0,
task.a().as_ptr(),
a[0],
a[1],
task.b().as_ptr(),
b[0],
b[1],
0.0,
product.as_mut_ptr(),
task.n() as isize,
1,
);
}
Some(product)
}
#[cfg(test)]
#[path = "tests/simd_tests.rs"]
mod tests;