use himada_core::HardwareDNA;
pub fn matmul_scalar(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
for i in 0..n {
for k in 0..n {
let aik = a[i * n + k];
let row_b = k * n;
let row_c = i * n;
for j in 0..n {
c[row_c + j] += aik * b[row_b + j];
}
}
}
}
pub fn matmul_scalar_supported(_: &HardwareDNA) -> bool {
true
}
pub fn matmul_tiled(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
const T: usize = 64;
for i in (0..n).step_by(T) {
let imax = (i + T).min(n);
for k in (0..n).step_by(T) {
let kmax = (k + T).min(n);
for j in (0..n).step_by(T) {
let jmax = (j + T).min(n);
for ii in i..imax {
for kk in k..kmax {
let aik = a[ii * n + kk];
let row_b = kk * n;
let row_c = ii * n;
for jj in j..jmax {
c[row_c + jj] += aik * b[row_b + jj];
}
}
}
}
}
}
}
pub fn matmul_tiled_supported(_: &HardwareDNA) -> bool {
true
}
#[cfg(target_arch = "x86_64")]
pub fn matmul_sse(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
unsafe {
for i in 0..n {
for k in 0..n {
let aik = a[i * n + k];
let row_b = k * n;
let row_c = i * n;
let mut j = 0;
if is_x86_feature_detected!("sse2") {
let vaik = _mm_set1_pd(aik);
while j + 2 <= n {
let vb = _mm_loadu_pd(b.as_ptr().add(row_b + j));
let vc = _mm_loadu_pd(c.as_ptr().add(row_c + j));
let vm = _mm_mul_pd(vaik, vb);
_mm_storeu_pd(c.as_mut_ptr().add(row_c + j), _mm_add_pd(vc, vm));
j += 2;
}
}
for jj in j..n {
c[row_c + jj] += aik * b[row_b + jj];
}
}
}
}
}
#[cfg(target_arch = "x86_64")]
pub fn matmul_sse_supported(dna: &HardwareDNA) -> bool {
dna.cpu.features.iter().any(|f| f == "SSE2")
}
#[cfg(not(target_arch = "x86_64"))]
pub fn matmul_sse(_: &[f64], _: &[f64], _: &mut [f64], _: usize) {}
#[cfg(not(target_arch = "x86_64"))]
pub fn matmul_sse_supported(_: &HardwareDNA) -> bool { false }
#[cfg(target_arch = "x86_64")]
pub fn matmul_avx2(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
unsafe {
for i in 0..n {
for k in 0..n {
let aik = a[i * n + k];
let row_b = k * n;
let row_c = i * n;
let mut j = 0;
if is_x86_feature_detected!("avx2") {
let vaik = _mm256_set1_pd(aik);
while j + 4 <= n {
let vb = _mm256_loadu_pd(b.as_ptr().add(row_b + j));
let vc = _mm256_loadu_pd(c.as_ptr().add(row_c + j));
let vm = _mm256_mul_pd(vaik, vb);
_mm256_storeu_pd(c.as_mut_ptr().add(row_c + j), _mm256_add_pd(vc, vm));
j += 4;
}
}
for jj in j..n {
c[row_c + jj] += aik * b[row_b + jj];
}
}
}
}
}
#[cfg(target_arch = "x86_64")]
pub fn matmul_avx2_supported(dna: &HardwareDNA) -> bool {
dna.cpu.features.iter().any(|f| f == "AVX2")
}
#[cfg(not(target_arch = "x86_64"))]
pub fn matmul_avx2(_: &[f64], _: &[f64], _: &mut [f64], _: usize) {}
#[cfg(not(target_arch = "x86_64"))]
pub fn matmul_avx2_supported(_: &HardwareDNA) -> bool { false }
#[cfg(target_arch = "aarch64")]
pub fn matmul_neon(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
#[cfg(target_arch = "aarch64")]
use std::arch::aarch64::*;
unsafe {
for i in 0..n {
for k in 0..n {
let aik = a[i * n + k];
let row_b = k * n;
let row_c = i * n;
let mut j = 0;
if n >= 2 {
let vaik = vdupq_n_f64(aik);
while j + 2 <= n {
let vb = vld1q_f64(b.as_ptr().add(row_b + j));
let vc = vld1q_f64(c.as_ptr().add(row_c + j));
let vm = vmulq_f64(vaik, vb);
vst1q_f64(c.as_mut_ptr().add(row_c + j), vaddq_f64(vc, vm));
j += 2;
}
}
for jj in j..n {
c[row_c + jj] += aik * b[row_b + jj];
}
}
}
}
}
#[cfg(target_arch = "aarch64")]
pub fn matmul_neon_supported(_: &HardwareDNA) -> bool {
true
}
#[cfg(not(target_arch = "aarch64"))]
pub fn matmul_neon(_: &[f64], _: &[f64], _: &mut [f64], _: usize) {}
#[cfg(not(target_arch = "aarch64"))]
pub fn matmul_neon_supported(_: &HardwareDNA) -> bool { false }
use himada_core::profile;
use std::sync::OnceLock;
fn optimal_tile_size_n() -> usize {
static TILE: OnceLock<usize> = OnceLock::new();
*TILE.get_or_init(|| {
let dna = profile::load_or_collect();
let l1 = dna.caches.l1_data_size.max(16384);
let elems = (l1 / (3 * 8)) as usize;
let mut t = 1;
while (t << 1) <= elems {
t <<= 1;
}
t.clamp(8, 256)
})
}
pub fn matmul_cache_tiled(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
let t = optimal_tile_size_n();
for i in (0..n).step_by(t) {
let imax = (i + t).min(n);
for k in (0..n).step_by(t) {
let kmax = (k + t).min(n);
for j in (0..n).step_by(t) {
let jmax = (j + t).min(n);
for ii in i..imax {
for kk in k..kmax {
let aik = a[ii * n + kk];
let row_b = kk * n;
let row_c = ii * n;
for jj in j..jmax {
c[row_c + jj] += aik * b[row_b + jj];
}
}
}
}
}
}
}
pub fn matmul_cache_tiled_supported(_: &HardwareDNA) -> bool { true }