use super::gemm_validate::{validate_gemm_bt, validate_gemm_nn};
#[cfg(not(target_os = "macos"))]
use super::simd::simd_config;
#[cfg(all(not(target_os = "macos"), target_arch = "aarch64"))]
use super::arch_kernels::matmul_neon;
#[cfg(all(not(target_os = "macos"), target_arch = "x86_64"))]
use super::arch_kernels::{matmul_avx2, matmul_avx512};
#[cfg(target_os = "macos")]
use super::blas::{accelerate_matmul, accelerate_matmul_bt};
#[cfg(not(target_os = "macos"))]
use super::tiled::matmul_bt_tiled;
pub fn matmul(a: &[f32], b: &[f32], m: usize, k: usize, n: usize) -> Vec<f32> {
assert!(
m.checked_mul(n).is_some(),
"matmul output shape overflow: m*n"
);
let mut c = vec![0.0f32; m * n];
matmul_into(a, b, &mut c, m, k, n);
c
}
pub fn matmul_into(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
validate_gemm_nn(a.len(), b.len(), c.len(), m, k, n, "matmul");
#[cfg(target_os = "macos")]
{
accelerate_matmul(a, b, c, m, n, k);
}
#[cfg(not(target_os = "macos"))]
matmul_scalar(a, b, c, m, k, n);
}
pub fn matmul_bt(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
validate_gemm_bt(a.len(), b.len(), c.len(), m, k, n, "matmul_bt");
#[cfg(target_os = "macos")]
{
accelerate_matmul_bt(a, b, c, m, n, k);
}
#[cfg(not(target_os = "macos"))]
{
let total_work = (m as u64) * (n as u64) * (k as u64);
if total_work >= 1024 * 1024 && k >= super::tiled::TILE_K {
matmul_bt_tiled(a, b, c, m, k, n);
return;
}
let config = simd_config();
#[cfg(target_arch = "x86_64")]
{
if config.avx512f_enabled && config.fma_enabled {
unsafe {
matmul_avx512(a, b, c, m, k, n);
return;
}
}
if config.avx2_enabled && config.fma_enabled {
unsafe {
matmul_avx2(a, b, c, m, k, n);
return;
}
}
}
#[cfg(target_arch = "aarch64")]
{
if config.neon_enabled {
unsafe {
matmul_neon(a, b, c, m, k, n);
return;
}
}
}
matmul_bt_scalar(a, b, c, m, k, n);
}
}
pub fn matmul_scalar(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
c.fill(0.0);
for i in 0..m {
for p in 0..k {
let a_val = a[i * k + p];
let b_row = &b[p * n..(p + 1) * n];
let c_row = &mut c[i * n..(i + 1) * n];
for j in 0..n {
c_row[j] += a_val * b_row[j];
}
}
}
}
#[cfg_attr(target_os = "macos", allow(dead_code))]
pub fn matmul_bt_scalar(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
if m == 1 {
matmul_bt_scalar_m1(a, b, c, k, n);
return;
}
for i in 0..m {
let a_row = &a[i * k..(i + 1) * k];
let c_row = &mut c[i * n..(i + 1) * n];
for j in 0..n {
let b_row = &b[j * k..(j + 1) * k];
let mut s0 = 0.0f32;
let mut s1 = 0.0f32;
let mut s2 = 0.0f32;
let mut s3 = 0.0f32;
let unrolled = k / 4;
for p in 0..unrolled {
let off = p * 4;
s0 += a_row[off] * b_row[off];
s1 += a_row[off + 1] * b_row[off + 1];
s2 += a_row[off + 2] * b_row[off + 2];
s3 += a_row[off + 3] * b_row[off + 3];
}
for p in (unrolled * 4)..k {
s0 += a_row[p] * b_row[p];
}
c_row[j] = (s0 + s1) + (s2 + s3);
}
}
}
#[cfg_attr(target_os = "macos", allow(dead_code))]
#[inline]
fn matmul_bt_scalar_m1(a: &[f32], b: &[f32], c: &mut [f32], k: usize, n: usize) {
let a_row = &a[..k];
let unrolled8 = k / 8;
for j in 0..n {
let b_row = &b[j * k..(j + 1) * k];
let mut s0 = 0.0f32;
let mut s1 = 0.0f32;
let mut s2 = 0.0f32;
let mut s3 = 0.0f32;
let mut s4 = 0.0f32;
let mut s5 = 0.0f32;
let mut s6 = 0.0f32;
let mut s7 = 0.0f32;
for p in 0..unrolled8 {
let off = p * 8;
s0 += a_row[off] * b_row[off];
s1 += a_row[off + 1] * b_row[off + 1];
s2 += a_row[off + 2] * b_row[off + 2];
s3 += a_row[off + 3] * b_row[off + 3];
s4 += a_row[off + 4] * b_row[off + 4];
s5 += a_row[off + 5] * b_row[off + 5];
s6 += a_row[off + 6] * b_row[off + 6];
s7 += a_row[off + 7] * b_row[off + 7];
}
for p in (unrolled8 * 8)..k {
s0 += a_row[p] * b_row[p];
}
c[j] = ((s0 + s1) + (s2 + s3)) + ((s4 + s5) + (s6 + s7));
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[should_panic(expected = "too short for n*k")]
fn matmul_bt_short_b_panics_in_release() {
let a = [0.0f32; 2]; let b: [f32; 0] = []; let mut c = [0.0f32; 2];
matmul_bt(&a, &b, &mut c, 1, 2, 2);
}
#[test]
#[should_panic(expected = "shape overflow")]
fn matmul_shape_overflow_panics() {
let a = [0.0f32; 2];
let b = [0.0f32; 2];
let mut c = [0.0f32; 1];
matmul_bt(&a, &b, &mut c, 2, usize::MAX, 2);
}
#[test]
fn matmul_bt_oversized_c_does_not_panic() {
let a = [1.0f32, 2.0];
let b = [1.0f32, 0.0, 0.0, 1.0]; let mut c = [0.0f32; 3]; matmul_bt(&a, &b, &mut c, 1, 2, 2);
assert!(
(c[0] - 1.0).abs() < 1e-6,
"c[0] should be 1.0, got {}",
c[0]
);
assert!(
(c[1] - 2.0).abs() < 1e-6,
"c[1] should be 2.0, got {}",
c[1]
);
}
}