#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
use rayon::prelude::*;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
const MR: usize = 6;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
const NR: usize = 16;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
const KC: usize = 256;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[derive(Clone, Copy)]
struct CPtr(*mut f32);
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe impl Send for CPtr {}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe impl Sync for CPtr {}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
impl CPtr {
#[inline]
fn get(self) -> *mut f32 {
self.0
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
pub(crate) fn sgemm_simd(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
if m == 0 || n == 0 {
return;
}
if k == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let m_panels = m.div_ceil(MR);
let mut apack = vec![0.0f32; m_panels * k * MR];
pack_a(a, &mut apack, m, k);
let n_panels = n.div_ceil(NR);
let threads = rayon::current_num_threads().max(1);
let target_tasks = threads.saturating_mul(8).max(1);
let panels_per_strip = n_panels.div_ceil(target_tasks).clamp(1, 16);
let strip_cols = panels_per_strip * NR;
let strip_count = n.div_ceil(strip_cols);
let cptr = CPtr(c.as_mut_ptr());
let apack = &apack;
(0..strip_count).into_par_iter().for_each(|s| {
let j0 = s * strip_cols;
let nc = strip_cols.min(n - j0);
let strip_panels = nc.div_ceil(NR);
let mut bpack = vec![0.0f32; KC * strip_panels * NR];
let c_base = cptr.get();
let mut pc = 0usize;
while pc < k {
let kc = KC.min(k - pc);
pack_b(b, &mut bpack, k, n, pc, kc, j0, nc);
let first = pc == 0;
for ip in 0..m_panels {
let i0 = ip * MR;
let mr = MR.min(m - i0);
let apanel = &apack[ip * k * MR + pc * MR..ip * k * MR + pc * MR + kc * MR];
let mut jr = 0usize;
let mut jp = 0usize;
while jr < nc {
let nr = NR.min(nc - jr);
let bpanel = &bpack[jp * KC * NR..jp * KC * NR + kc * NR];
unsafe {
micro_6x16(
apanel.as_ptr(),
bpanel.as_ptr(),
c_base.add(i0 * n + j0 + jr),
n,
kc,
mr,
nr,
first,
);
}
jr += NR;
jp += 1;
}
}
pc += KC;
}
});
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
fn pack_a(a: &[f32], apack: &mut [f32], m: usize, k: usize) {
let m_panels = m.div_ceil(MR);
for ip in 0..m_panels {
let i0 = ip * MR;
let mr = MR.min(m - i0);
let dst = &mut apack[ip * k * MR..ip * k * MR + k * MR];
for p in 0..k {
let out = &mut dst[p * MR..p * MR + MR];
for r in 0..mr {
out[r] = a[(i0 + r) * k + p];
}
}
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[allow(clippy::too_many_arguments)]
fn pack_b(
b: &[f32],
bpack: &mut [f32],
_k: usize,
n: usize,
pc: usize,
kc: usize,
j0: usize,
nc: usize,
) {
let n_panels = nc.div_ceil(NR);
for jp in 0..n_panels {
let jcol = j0 + jp * NR;
let nr = NR.min(nc - jp * NR);
let dst = &mut bpack[jp * KC * NR..jp * KC * NR + kc * NR];
for p in 0..kc {
let src = &b[(pc + p) * n + jcol..(pc + p) * n + jcol + nr];
let out = &mut dst[p * NR..p * NR + NR];
out[..nr].copy_from_slice(src);
out[nr..NR].fill(0.0);
}
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "avx2,fma")]
#[allow(clippy::too_many_arguments)]
unsafe fn micro_6x16(
apack: *const f32,
bpack: *const f32,
c: *mut f32,
n: usize,
kc: usize,
mr: usize,
nr: usize,
first: bool,
) {
#[cfg(target_arch = "x86")]
use std::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
unsafe {
let mut c0 = [_mm256_setzero_ps(); MR];
let mut c1 = [_mm256_setzero_ps(); MR];
let mut p = 0usize;
while p < kc {
let b0 = _mm256_loadu_ps(bpack.add(p * NR));
let b1 = _mm256_loadu_ps(bpack.add(p * NR + 8));
let arow = apack.add(p * MR);
let a0 = _mm256_broadcast_ss(&*arow.add(0));
c0[0] = _mm256_fmadd_ps(a0, b0, c0[0]);
c1[0] = _mm256_fmadd_ps(a0, b1, c1[0]);
let a1 = _mm256_broadcast_ss(&*arow.add(1));
c0[1] = _mm256_fmadd_ps(a1, b0, c0[1]);
c1[1] = _mm256_fmadd_ps(a1, b1, c1[1]);
let a2 = _mm256_broadcast_ss(&*arow.add(2));
c0[2] = _mm256_fmadd_ps(a2, b0, c0[2]);
c1[2] = _mm256_fmadd_ps(a2, b1, c1[2]);
let a3 = _mm256_broadcast_ss(&*arow.add(3));
c0[3] = _mm256_fmadd_ps(a3, b0, c0[3]);
c1[3] = _mm256_fmadd_ps(a3, b1, c1[3]);
let a4 = _mm256_broadcast_ss(&*arow.add(4));
c0[4] = _mm256_fmadd_ps(a4, b0, c0[4]);
c1[4] = _mm256_fmadd_ps(a4, b1, c1[4]);
let a5 = _mm256_broadcast_ss(&*arow.add(5));
c0[5] = _mm256_fmadd_ps(a5, b0, c0[5]);
c1[5] = _mm256_fmadd_ps(a5, b1, c1[5]);
p += 1;
}
if nr == NR {
for r in 0..mr {
let dst = c.add(r * n);
if first {
_mm256_storeu_ps(dst, c0[r]);
_mm256_storeu_ps(dst.add(8), c1[r]);
} else {
let old0 = _mm256_loadu_ps(dst);
let old1 = _mm256_loadu_ps(dst.add(8));
_mm256_storeu_ps(dst, _mm256_add_ps(old0, c0[r]));
_mm256_storeu_ps(dst.add(8), _mm256_add_ps(old1, c1[r]));
}
}
} else {
let mut tmp = [0.0f32; NR];
for r in 0..mr {
_mm256_storeu_ps(tmp.as_mut_ptr(), c0[r]);
_mm256_storeu_ps(tmp.as_mut_ptr().add(8), c1[r]);
let dst = c.add(r * n);
for (col, &val) in tmp[..nr].iter().enumerate() {
if first {
*dst.add(col) = val;
} else {
*dst.add(col) += val;
}
}
}
}
}
}
#[cfg(all(test, any(target_arch = "x86", target_arch = "x86_64")))]
mod tests {
use super::*;
use crate::backend::has_simd_x86;
fn reference(a: &[f32], b: &[f32], m: usize, k: usize, n: usize) -> Vec<f32> {
let mut c = vec![0.0f32; m * n];
for i in 0..m {
for p in 0..k {
let aip = a[i * k + p];
for j in 0..n {
c[i * n + j] += aip * b[p * n + j];
}
}
}
c
}
fn fill(len: usize, seed: usize) -> Vec<f32> {
(0..len)
.map(|i| (((i + seed) as f32 * 0.123).sin()) * 2.0 - 0.5)
.collect()
}
fn check(m: usize, k: usize, n: usize) {
if !has_simd_x86() {
return; }
let a = fill(m * k, 1);
let b = fill(k * n, 7);
let expect = reference(&a, &b, m, k, n);
let mut got = vec![0.0f32; m * n];
sgemm_simd(&a, &b, &mut got, m, k, n);
for (idx, (g, e)) in got.iter().zip(expect.iter()).enumerate() {
let tol = 1e-3 * (1.0 + e.abs());
assert!(
(g - e).abs() <= tol,
"mismatch at {idx} for {m}x{k}x{n}: got {g}, expect {e}"
);
}
}
#[test]
fn exact_tile_multiple() {
check(12, 64, 32);
}
#[test]
fn tail_shapes() {
check(7, 33, 17);
check(1, 5, 3);
check(6, 16, 16);
check(5, 1, 5);
}
#[test]
fn thin_vectors() {
check(1, 128, 1); check(1, 512, 256); check(256, 512, 1); }
#[test]
fn multi_kpanel() {
check(9, KC * 2 + 13, 40);
}
#[test]
fn zero_dims() {
let mut c = vec![1.0f32; 4];
sgemm_simd(&[], &[], &mut c, 0, 3, 4); sgemm_simd(&[1.0], &[], &mut c, 2, 0, 2); assert_eq!(&c[..4], &[0.0, 0.0, 0.0, 0.0]);
}
}