#[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> {
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) {
debug_assert_eq!(a.len(), m * k);
debug_assert_eq!(b.len(), k * n);
debug_assert_eq!(c.len(), m * n);
#[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) {
debug_assert_eq!(a.len(), m * k);
debug_assert_eq!(b.len(), n * k);
debug_assert_eq!(c.len(), m * n);
#[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));
}
}