#[cfg(not(target_os = "macos"))]
use super::matmul::{matmul_bt, matmul_into};
#[cfg(target_os = "macos")]
mod 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,
);
}
pub const CBLAS_ROW_MAJOR: i32 = 101;
pub const CBLAS_NO_TRANS: i32 = 111;
pub const CBLAS_TRANS: i32 = 112;
}
#[cfg(target_os = "macos")]
pub(super) fn accelerate_matmul_bt(
a: &[f32],
b: &[f32],
output: &mut [f32],
m: usize,
n: usize,
k: usize,
) {
assert!(a.len() >= m * k, "A too small: {} < {}", a.len(), m * k);
assert!(b.len() >= n * k, "B too small: {} < {}", b.len(), n * k);
assert!(
output.len() >= m * n,
"C too small: {} < {}",
output.len(),
m * n
);
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_TRANS,
m as i32,
n as i32,
k as i32,
1.0,
a.as_ptr(),
k as i32,
b.as_ptr(),
k as i32,
0.0,
output.as_mut_ptr(),
n as i32,
);
}
}
#[cfg(target_os = "macos")]
pub(super) fn accelerate_matmul(
a: &[f32],
b: &[f32],
output: &mut [f32],
m: usize,
n: usize,
k: usize,
) {
assert!(a.len() >= m * k, "A too small: {} < {}", a.len(), m * k);
assert!(b.len() >= k * n, "B too small: {} < {}", b.len(), k * n);
assert!(
output.len() >= m * n,
"C too small: {} < {}",
output.len(),
m * n
);
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_NO_TRANS,
m as i32,
n as i32,
k as i32,
1.0,
a.as_ptr(),
k as i32,
b.as_ptr(),
n as i32,
0.0,
output.as_mut_ptr(),
n as i32,
);
}
}
#[cfg(target_os = "macos")]
pub unsafe fn sgemm_bt_strided(
a: *const f32,
lda: usize,
b: *const f32,
ldb: usize,
c: *mut f32,
ldc: usize,
m: usize,
n: usize,
k: usize,
) {
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_TRANS,
m as i32,
n as i32,
k as i32,
1.0,
a,
lda as i32,
b,
ldb as i32,
0.0,
c,
ldc as i32,
);
}
}
#[cfg(target_os = "macos")]
pub unsafe fn sgemm_nn_strided(
a: *const f32,
lda: usize,
b: *const f32,
ldb: usize,
c: *mut f32,
ldc: usize,
m: usize,
n: usize,
k: usize,
) {
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_NO_TRANS,
m as i32,
n as i32,
k as i32,
1.0,
a,
lda as i32,
b,
ldb as i32,
0.0,
c,
ldc as i32,
);
}
}
#[cfg(target_os = "macos")]
pub fn sgemm_nn_ab(
a: &[f32],
b: &[f32],
c: &mut [f32],
m: usize,
n: usize,
k: usize,
alpha: f32,
beta: f32,
) {
debug_assert!(a.len() >= m * k, "A too small: {} < {}", a.len(), m * k);
debug_assert!(b.len() >= k * n, "B too small: {} < {}", b.len(), k * n);
debug_assert!(c.len() >= m * n, "C too small: {} < {}", c.len(), m * n);
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_NO_TRANS,
m as i32,
n as i32,
k as i32,
alpha,
a.as_ptr(),
k as i32,
b.as_ptr(),
n as i32,
beta,
c.as_mut_ptr(),
n as i32,
);
}
}
#[cfg(not(target_os = "macos"))]
pub fn sgemm_nn_ab(
a: &[f32],
b: &[f32],
c: &mut [f32],
m: usize,
n: usize,
k: usize,
alpha: f32,
beta: f32,
) {
debug_assert!(a.len() >= m * k);
debug_assert!(b.len() >= k * n);
debug_assert!(c.len() >= m * n);
for i in 0..m {
for j in 0..n {
let mut sum = 0.0f32;
for p in 0..k {
sum += a[i * k + p] * b[p * n + j];
}
c[i * n + j] = alpha * sum + beta * c[i * n + j];
}
}
}
#[cfg(not(target_os = "macos"))]
pub unsafe fn sgemm_bt_strided(
a: *const f32,
lda: usize,
b: *const f32,
ldb: usize,
c: *mut f32,
ldc: usize,
m: usize,
n: usize,
k: usize,
) {
let mut a_contig = vec![0.0f32; m * k];
for row in 0..m {
let src = unsafe { std::slice::from_raw_parts(a.add(row * lda), k) };
a_contig[row * k..(row + 1) * k].copy_from_slice(src);
}
let mut b_contig = vec![0.0f32; n * k];
for row in 0..n {
let src = unsafe { std::slice::from_raw_parts(b.add(row * ldb), k) };
b_contig[row * k..(row + 1) * k].copy_from_slice(src);
}
let mut c_contig = vec![0.0f32; m * n];
matmul_bt(&a_contig, &b_contig, &mut c_contig, m, k, n);
for row in 0..m {
let dst = unsafe { std::slice::from_raw_parts_mut(c.add(row * ldc), n) };
dst.copy_from_slice(&c_contig[row * n..(row + 1) * n]);
}
}
#[cfg(not(target_os = "macos"))]
pub unsafe fn sgemm_nn_strided(
a: *const f32,
lda: usize,
b: *const f32,
ldb: usize,
c: *mut f32,
ldc: usize,
m: usize,
n: usize,
k: usize,
) {
let mut a_contig = vec![0.0f32; m * k];
for row in 0..m {
let src = unsafe { std::slice::from_raw_parts(a.add(row * lda), k) };
a_contig[row * k..(row + 1) * k].copy_from_slice(src);
}
let mut b_contig = vec![0.0f32; k * n];
for row in 0..k {
let src = unsafe { std::slice::from_raw_parts(b.add(row * ldb), n) };
b_contig[row * n..(row + 1) * n].copy_from_slice(src);
}
let mut c_contig = vec![0.0f32; m * n];
matmul_into(&a_contig, &b_contig, &mut c_contig, m, k, n);
for row in 0..m {
let dst = unsafe { std::slice::from_raw_parts_mut(c.add(row * ldc), n) };
dst.copy_from_slice(&c_contig[row * n..(row + 1) * n]);
}
}