#[cfg(target_os = "macos")]
use super::gemm_validate::validate_gemm_bt;
use super::gemm_validate::{validate_gemm_nn, validate_gemm_strided_shape};
#[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")]
#[inline]
fn cblas_dim(value: usize, param: &'static str, op: &'static str) -> i32 {
i32::try_from(value)
.unwrap_or_else(|_| panic!("{op}: {param}={value} exceeds i32::MAX (CBLAS ABI limit)"))
}
#[cfg(target_os = "macos")]
pub(super) fn accelerate_matmul_bt(
a: &[f32],
b: &[f32],
output: &mut [f32],
m: usize,
n: usize,
k: usize,
) {
validate_gemm_bt(
a.len(),
b.len(),
output.len(),
m,
k,
n,
"accelerate_matmul_bt",
);
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_TRANS,
cblas_dim(m, "m", "accelerate_matmul_bt"),
cblas_dim(n, "n", "accelerate_matmul_bt"),
cblas_dim(k, "k", "accelerate_matmul_bt"),
1.0,
a.as_ptr(),
cblas_dim(k, "lda", "accelerate_matmul_bt"),
b.as_ptr(),
cblas_dim(k, "ldb", "accelerate_matmul_bt"),
0.0,
output.as_mut_ptr(),
cblas_dim(n, "ldc", "accelerate_matmul_bt"),
);
}
}
#[cfg(target_os = "macos")]
pub(super) fn accelerate_matmul(
a: &[f32],
b: &[f32],
output: &mut [f32],
m: usize,
n: usize,
k: usize,
) {
validate_gemm_nn(a.len(), b.len(), output.len(), m, k, n, "accelerate_matmul");
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_NO_TRANS,
cblas_dim(m, "m", "accelerate_matmul"),
cblas_dim(n, "n", "accelerate_matmul"),
cblas_dim(k, "k", "accelerate_matmul"),
1.0,
a.as_ptr(),
cblas_dim(k, "lda", "accelerate_matmul"),
b.as_ptr(),
cblas_dim(n, "ldb", "accelerate_matmul"),
0.0,
output.as_mut_ptr(),
cblas_dim(n, "ldc", "accelerate_matmul"),
);
}
}
#[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,
) {
validate_gemm_strided_shape(m, k, n, lda, ldb, ldc, true, "sgemm_bt_strided");
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_TRANS,
cblas_dim(m, "m", "sgemm_bt_strided"),
cblas_dim(n, "n", "sgemm_bt_strided"),
cblas_dim(k, "k", "sgemm_bt_strided"),
1.0,
a,
cblas_dim(lda, "lda", "sgemm_bt_strided"),
b,
cblas_dim(ldb, "ldb", "sgemm_bt_strided"),
0.0,
c,
cblas_dim(ldc, "ldc", "sgemm_bt_strided"),
);
}
}
#[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,
) {
validate_gemm_strided_shape(m, k, n, lda, ldb, ldc, false, "sgemm_nn_strided");
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_NO_TRANS,
cblas_dim(m, "m", "sgemm_nn_strided"),
cblas_dim(n, "n", "sgemm_nn_strided"),
cblas_dim(k, "k", "sgemm_nn_strided"),
1.0,
a,
cblas_dim(lda, "lda", "sgemm_nn_strided"),
b,
cblas_dim(ldb, "ldb", "sgemm_nn_strided"),
0.0,
c,
cblas_dim(ldc, "ldc", "sgemm_nn_strided"),
);
}
}
#[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,
) {
validate_gemm_nn(a.len(), b.len(), c.len(), m, k, n, "sgemm_nn_ab");
unsafe {
accelerate::cblas_sgemm(
accelerate::CBLAS_ROW_MAJOR,
accelerate::CBLAS_NO_TRANS,
accelerate::CBLAS_NO_TRANS,
cblas_dim(m, "m", "sgemm_nn_ab"),
cblas_dim(n, "n", "sgemm_nn_ab"),
cblas_dim(k, "k", "sgemm_nn_ab"),
alpha,
a.as_ptr(),
cblas_dim(k, "lda", "sgemm_nn_ab"),
b.as_ptr(),
cblas_dim(n, "ldb", "sgemm_nn_ab"),
beta,
c.as_mut_ptr(),
cblas_dim(n, "ldc", "sgemm_nn_ab"),
);
}
}
#[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,
) {
validate_gemm_nn(a.len(), b.len(), c.len(), m, k, n, "sgemm_nn_ab");
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,
) {
validate_gemm_strided_shape(m, k, n, lda, ldb, ldc, true, "sgemm_bt_strided");
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,
) {
validate_gemm_strided_shape(m, k, n, lda, ldb, ldc, false, "sgemm_nn_strided");
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]);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[should_panic(expected = "a too short for m*k")]
fn sgemm_nn_ab_rejects_short_a() {
let a = [0.0f32; 1]; let b = [0.0f32; 2];
let mut c = [0.0f32; 1];
sgemm_nn_ab(&a, &b, &mut c, 1, 1, 2, 1.0, 0.0);
}
#[test]
#[should_panic(expected = "b too short for k*n")]
fn sgemm_nn_ab_rejects_short_b() {
let a = [0.0f32; 2];
let b = [0.0f32; 1]; let mut c = [0.0f32; 1];
sgemm_nn_ab(&a, &b, &mut c, 1, 2, 1, 1.0, 0.0);
}
#[test]
#[should_panic(expected = "shape overflow: m*k")]
fn sgemm_nn_ab_rejects_overflow() {
let a = [0.0f32; 2];
let b = [0.0f32; 2];
let mut c = [0.0f32; 2];
sgemm_nn_ab(&a, &b, &mut c, usize::MAX, 2, 2, 1.0, 0.0);
}
#[test]
fn sgemm_nn_ab_accepts_oversized_buffers_and_computes_correctly() {
let a = [1.0f32, 2.0, 99.0]; let b = [1.0f32, 0.0, 0.0, 1.0, 99.0]; let mut c = [0.0f32, 0.0, 99.0]; sgemm_nn_ab(&a, &b, &mut c, 1, 2, 2, 1.0, 0.0);
assert!((c[0] - 1.0).abs() < 1e-6);
assert!((c[1] - 2.0).abs() < 1e-6);
}
#[cfg(target_os = "macos")]
#[test]
#[should_panic(expected = "b too short for n*k")]
fn accelerate_matmul_bt_rejects_short_b() {
let a = [0.0f32; 2];
let b = [0.0f32; 1]; let mut c = [0.0f32; 1];
accelerate_matmul_bt(&a, &b, &mut c, 1, 1, 2);
}
#[cfg(target_os = "macos")]
#[test]
#[should_panic(expected = "shape overflow: n*k")]
fn accelerate_matmul_bt_rejects_overflow() {
let a = [0.0f32; 2];
let b = [0.0f32; 2];
let mut c = [0.0f32; 2];
accelerate_matmul_bt(&a, &b, &mut c, 2, usize::MAX, 2);
}
#[cfg(target_os = "macos")]
#[test]
#[should_panic(expected = "b too short for k*n")]
fn accelerate_matmul_rejects_short_b() {
let a = [0.0f32; 2];
let b = [0.0f32; 1]; let mut c = [0.0f32; 1];
accelerate_matmul(&a, &b, &mut c, 1, 1, 2);
}
#[test]
#[should_panic(expected = "ldb too small for row extent n")]
fn sgemm_nn_strided_rejects_short_ldb() {
let a = [0.0f32; 4];
let b = [0.0f32; 4];
let mut c = [0.0f32; 4];
unsafe {
sgemm_nn_strided(a.as_ptr(), 2, b.as_ptr(), 1, c.as_mut_ptr(), 2, 1, 2, 2);
}
}
#[test]
#[should_panic(expected = "ldb too small for row extent k")]
fn sgemm_bt_strided_rejects_short_ldb() {
let a = [0.0f32; 4];
let b = [0.0f32; 4];
let mut c = [0.0f32; 4];
unsafe {
sgemm_bt_strided(a.as_ptr(), 2, b.as_ptr(), 1, c.as_mut_ptr(), 2, 1, 2, 2);
}
}
#[cfg(target_os = "macos")]
#[test]
#[should_panic(expected = "n=2147483648 exceeds i32::MAX")]
fn accelerate_matmul_rejects_n_above_i32_max() {
let huge_n = i32::MAX as usize + 1;
let a: [f32; 0] = [];
let b: [f32; 0] = [];
let mut c: [f32; 0] = [];
accelerate_matmul(&a, &b, &mut c, 0, huge_n, 0);
}
#[cfg(target_os = "macos")]
#[test]
#[should_panic(expected = "n=2147483648 exceeds i32::MAX")]
fn accelerate_matmul_bt_rejects_n_above_i32_max() {
let huge_n = i32::MAX as usize + 1;
let a: [f32; 0] = [];
let b: [f32; 0] = [];
let mut c: [f32; 0] = [];
accelerate_matmul_bt(&a, &b, &mut c, 0, huge_n, 0);
}
#[cfg(target_os = "macos")]
#[test]
#[should_panic(expected = "ldb=2147483648 exceeds i32::MAX")]
fn sgemm_nn_strided_rejects_ldb_above_i32_max() {
let huge_ldb = i32::MAX as usize + 1;
let a: [f32; 0] = [];
let b: [f32; 0] = [];
let mut c: [f32; 0] = [];
unsafe {
sgemm_nn_strided(
a.as_ptr(),
0,
b.as_ptr(),
huge_ldb,
c.as_mut_ptr(),
0,
0,
0,
0,
);
}
}
#[cfg(target_os = "macos")]
#[test]
#[should_panic(expected = "m=2147483648 exceeds i32::MAX")]
fn sgemm_nn_ab_rejects_m_above_i32_max() {
let huge_m = i32::MAX as usize + 1;
let a: [f32; 0] = [];
let b: [f32; 0] = [];
let mut c: [f32; 0] = [];
sgemm_nn_ab(&a, &b, &mut c, huge_m, 0, 0, 1.0, 0.0);
}
}