#[cfg(feature = "blas")]
mod blas_backend {
#[cfg_attr(feature = "accelerate", link(name = "Accelerate", kind = "framework"))]
unsafe extern "C" {
pub fn cblas_sgemm(
order: i32,
transa: i32,
transb: i32,
m: i32,
n: i32,
k: i32,
alpha: f32,
a: *const f32,
lda: i32,
b: *const f32,
ldb: i32,
beta: f32,
c: *mut f32,
ldc: i32,
);
}
}
#[inline]
pub fn sgemm_ld(
trans_a: bool,
trans_b: bool,
m: usize,
k: usize,
n: usize,
a: &[f32],
lda: usize,
b: &[f32],
ldb: usize,
c: &mut [f32],
ldc: usize,
) {
#[cfg(feature = "blas")]
unsafe {
blas_backend::cblas_sgemm(
101,
if trans_a { 112 } else { 111 },
if trans_b { 112 } else { 111 },
m as i32,
n as i32,
k as i32,
1.0,
a.as_ptr(),
lda as i32,
b.as_ptr(),
ldb as i32,
0.0,
c.as_mut_ptr(),
ldc as i32,
);
}
#[cfg(not(feature = "blas"))]
unsafe {
let (rsa, csa) = if trans_a {
(1isize, lda as isize)
} else {
(lda as isize, 1)
};
let (rsb, csb) = if trans_b {
(1isize, ldb as isize)
} else {
(ldb as isize, 1)
};
matrixmultiply::sgemm(
m,
k,
n,
1.0,
a.as_ptr(),
rsa,
csa,
b.as_ptr(),
rsb,
csb,
0.0,
c.as_mut_ptr(),
ldc as isize,
1,
);
}
}
pub fn sgemm_row_major(m: usize, k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
debug_assert!(a.len() >= m * k);
debug_assert!(b.len() >= k * n);
debug_assert!(c.len() >= m * n);
sgemm_ld(false, false, m, k, n, a, k, b, n, c, n);
}
pub fn sgemm_row_major_b_transposed(
m: usize,
k: usize,
n: usize,
a: &[f32],
b: &[f32],
c: &mut [f32],
) {
debug_assert!(a.len() >= m * k);
debug_assert!(b.len() >= n * k);
debug_assert!(c.len() >= m * n);
sgemm_ld(false, true, m, k, n, a, k, b, k, c, n);
}