use kime_tensor::Epilogue;
use crate::ops::gelu;
use crate::par::{self, Shared};
const LANES: usize = 8;
const BLOCK: usize = 64;
const MR: usize = 4;
const NR: usize = 3;
const NB: usize = 48;
const MB: usize = 128;
#[must_use]
pub fn dot(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len());
#[cfg(target_arch = "aarch64")]
return tile::<neon::Neon, 1, 1>([a], [b])[0][0];
#[cfg(target_arch = "x86_64")]
if has_fma() {
return unsafe { dot_fma(a, b) };
}
#[allow(unreachable_code)]
tile::<[f32; LANES], 1, 1>([a], [b])[0][0]
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
fn dot_fma(a: &[f32], b: &[f32]) -> f32 {
tile::<avx::Avx, 1, 1>([a], [b])[0][0]
}
#[cfg(target_arch = "x86_64")]
#[inline]
fn has_fma() -> bool {
std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma")
}
trait V8: Copy {
fn zero() -> Self;
fn load(s: &[f32; LANES]) -> Self;
fn fma(self, a: Self, b: Self) -> Self;
fn lanes(self) -> [f32; LANES];
}
impl V8 for [f32; LANES] {
#[inline(always)]
fn zero() -> Self {
[0.0; LANES]
}
#[inline(always)]
fn load(s: &[f32; LANES]) -> Self {
*s
}
#[inline(always)]
fn fma(self, a: Self, b: Self) -> Self {
std::array::from_fn(|l| a[l].mul_add(b[l], self[l]))
}
#[inline(always)]
fn lanes(self) -> [f32; LANES] {
self
}
}
#[cfg(target_arch = "aarch64")]
mod neon {
use std::arch::aarch64::{float32x4_t, vdupq_n_f32, vfmaq_f32, vld1q_f32, vst1q_f32};
use super::{LANES, V8};
#[derive(Clone, Copy)]
pub(super) struct Neon(float32x4_t, float32x4_t);
impl V8 for Neon {
#[inline(always)]
fn zero() -> Self {
unsafe { Self(vdupq_n_f32(0.0), vdupq_n_f32(0.0)) }
}
#[inline(always)]
fn load(s: &[f32; LANES]) -> Self {
unsafe { Self(vld1q_f32(s.as_ptr()), vld1q_f32(s.as_ptr().add(4))) }
}
#[inline(always)]
fn fma(self, a: Self, b: Self) -> Self {
unsafe { Self(vfmaq_f32(self.0, a.0, b.0), vfmaq_f32(self.1, a.1, b.1)) }
}
#[inline(always)]
fn lanes(self) -> [f32; LANES] {
let mut out = [0f32; LANES];
unsafe {
vst1q_f32(out.as_mut_ptr(), self.0);
vst1q_f32(out.as_mut_ptr().add(4), self.1);
}
out
}
}
}
#[cfg(target_arch = "x86_64")]
mod avx {
use std::arch::x86_64::{
__m256, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_setzero_ps, _mm256_storeu_ps,
};
use super::{LANES, V8};
#[derive(Clone, Copy)]
pub(super) struct Avx(__m256);
impl V8 for Avx {
#[inline(always)]
fn zero() -> Self {
unsafe { Self(_mm256_setzero_ps()) }
}
#[inline(always)]
fn load(s: &[f32; LANES]) -> Self {
unsafe { Self(_mm256_loadu_ps(s.as_ptr())) }
}
#[inline(always)]
fn fma(self, a: Self, b: Self) -> Self {
unsafe { Self(_mm256_fmadd_ps(a.0, b.0, self.0)) }
}
#[inline(always)]
fn lanes(self) -> [f32; LANES] {
let mut out = [0f32; LANES];
unsafe { _mm256_storeu_ps(out.as_mut_ptr(), self.0) };
out
}
}
}
#[inline(always)]
fn tile<V: V8, const R: usize, const C: usize>(x: [&[f32]; R], w: [&[f32]; C]) -> [[f32; C]; R] {
let k = w[0].len();
let body = k - k % LANES;
let mut wide = [[[0f64; LANES]; C]; R];
let mut p = 0;
while p < body {
let end = (p + BLOCK).min(body);
let mut acc = [[V::zero(); C]; R];
while p < end {
let xv: [V; R] =
std::array::from_fn(|r| V::load(x[r][p..p + LANES].try_into().unwrap()));
let wv: [V; C] =
std::array::from_fn(|c| V::load(w[c][p..p + LANES].try_into().unwrap()));
for r in 0..R {
for c in 0..C {
acc[r][c] = acc[r][c].fma(xv[r], wv[c]);
}
}
p += LANES;
}
for r in 0..R {
for c in 0..C {
let l = acc[r][c].lanes();
for i in 0..LANES {
wide[r][c][i] += f64::from(l[i]);
}
}
}
}
let mut out = [[0f32; C]; R];
for r in 0..R {
for c in 0..C {
let v = wide[r][c];
let mut s = ((v[0] + v[4]) + (v[2] + v[6])) + ((v[1] + v[5]) + (v[3] + v[7]));
for q in body..k {
s = f64::from(x[r][q]).mul_add(f64::from(w[c][q]), s);
}
out[r][c] = s as f32;
}
}
out
}
struct Args<'a> {
x: &'a [f32],
w: &'a [f32],
b: Option<&'a [f32]>,
ep: Epilogue,
k: usize,
n: usize,
y: &'a Shared<'a>,
}
impl Args<'_> {
#[inline(always)]
fn run<V: V8, const R: usize, const C: usize>(&self, i: usize, j: usize) {
let k = self.k;
let xs = std::array::from_fn(|r| &self.x[(i + r) * k..(i + r + 1) * k]);
let ws = std::array::from_fn(|c| &self.w[(j + c) * k..(j + c + 1) * k]);
let out = tile::<V, R, C>(xs, ws);
for (r, row) in out.iter().enumerate() {
for (c, &v) in row.iter().enumerate() {
self.put(i + r, j + c, v);
}
}
}
#[inline(always)]
fn put(&self, i: usize, j: usize, v: f32) {
let v = match self.b {
Some(b) => v + b[j],
None => v,
};
let at = i * self.n + j;
let v = match self.ep {
Epilogue::None => v,
Epilogue::Gelu => gelu(v),
Epilogue::Relu => v.max(0.0),
Epilogue::Accumulate => v + unsafe { self.y.get(at) },
};
unsafe { self.y.set(at, v) };
}
#[inline(always)]
fn block<V: V8>(&self, rows: (usize, usize), cols: (usize, usize)) {
let mut j = cols.0;
while j < cols.1 {
let wide = cols.1 - j >= NR;
let mut i = rows.0;
while i + MR <= rows.1 {
if wide {
self.run::<V, MR, NR>(i, j);
} else {
for c in j..cols.1 {
self.run::<V, MR, 1>(i, c);
}
}
i += MR;
}
for r in i..rows.1 {
if wide {
self.run::<V, 1, NR>(r, j);
} else {
for c in j..cols.1 {
self.run::<V, 1, 1>(r, c);
}
}
}
j += NR;
}
}
fn block_dispatch(&self, rows: (usize, usize), cols: (usize, usize)) {
#[cfg(target_arch = "aarch64")]
return self.block::<neon::Neon>(rows, cols);
#[cfg(target_arch = "x86_64")]
if has_fma() {
unsafe { self.block_fma(rows, cols) };
return;
}
#[allow(unreachable_code)]
self.block::<[f32; LANES]>(rows, cols);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
fn block_fma(&self, rows: (usize, usize), cols: (usize, usize)) {
self.block::<avx::Avx>(rows, cols);
}
}
#[allow(clippy::too_many_arguments)]
pub fn linear(
x: &[f32],
m: usize,
k: usize,
w: &[f32],
n: usize,
b: Option<&[f32]>,
y: &mut [f32],
threads: usize,
) {
let g = Gemm { x, m, k, w, n, b, ep: Epilogue::None };
g.run(y, threads, |tasks, f| par::for_each(tasks, threads, f));
}
#[derive(Debug, Clone, Copy)]
pub struct Gemm<'a> {
pub x: &'a [f32],
pub m: usize,
pub k: usize,
pub w: &'a [f32],
pub n: usize,
pub b: Option<&'a [f32]>,
pub ep: Epilogue,
}
impl Gemm<'_> {
fn row_block(&self, threads: usize) -> usize {
let nt = self.n.div_ceil(NB);
[MB, 64, 32, 16]
.into_iter()
.find(|&mb| self.m.div_ceil(mb) * nt >= 3 * threads)
.unwrap_or(16)
}
pub fn run(
&self,
y: &mut [f32],
threads: usize,
spawn: impl FnOnce(usize, &(dyn Fn(usize) + Sync)),
) {
let Self { x, m, k, w, n, b, ep } = *self;
assert_eq!(x.len(), m * k, "x is not [m, k]");
assert_eq!(w.len(), n * k, "w is not [n, k]");
assert_eq!(y.len(), m * n, "y is not [m, n]");
if let Some(b) = b {
assert_eq!(b.len(), n, "b is not [n]");
}
if m == 0 || n == 0 {
return;
}
let shared = Shared::new(y);
let args = Args { x, w, b, ep, k, n, y: &shared };
if k == 0 {
for i in 0..m {
for j in 0..n {
args.put(i, j, 0.0);
}
}
return;
}
let mb = self.row_block(threads);
let mt = m.div_ceil(mb);
let nt = n.div_ceil(NB);
spawn(mt * nt, &|t| {
let (bi, bj) = (t % mt, t / mt);
let rows = (bi * mb, ((bi + 1) * mb).min(m));
let cols = (bj * NB, ((bj + 1) * NB).min(n));
args.block_dispatch(rows, cols);
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::testing::{Rng, close};
fn naive(x: &[f32], m: usize, k: usize, w: &[f32], n: usize, b: Option<&[f32]>) -> Vec<f32> {
let mut y = vec![0f32; m * n];
for i in 0..m {
for j in 0..n {
let s: f64 =
(0..k).map(|p| f64::from(x[i * k + p]) * f64::from(w[j * k + p])).sum();
y[i * n + j] = (s + b.map_or(0.0, |b| f64::from(b[j]))) as f32;
}
}
y
}
#[test]
fn matches_naive_on_awkward_shapes() {
let mut rng = Rng(7);
let shapes = [
(0, 8, 5),
(1, 1, 1),
(1, 7, 3),
(3, 16, 2),
(4, 9, 3),
(5, 64, 7),
(130, 33, 50),
(129, 1028, 49),
(17, 0, 4),
];
for (m, k, n) in shapes {
let x = rng.vec(m * k);
let w = rng.vec(n * k);
let b = rng.vec(n);
for bias in [None, Some(&b[..])] {
let want = naive(&x, m, k, &w, n, bias);
for threads in [1, 3] {
let mut y = vec![f32::NAN; m * n];
linear(&x, m, k, &w, n, bias, &mut y, threads);
close(&y, &want, 1e-5, &format!("{m}x{k}x{n}"));
}
}
}
}
#[test]
fn same_bits_for_any_split() {
let mut rng = Rng(11);
let (m, k, n) = (37, 200, 101);
let x = rng.vec(m * k);
let w = rng.vec(n * k);
let mut one = vec![0f32; m * n];
linear(&x, m, k, &w, n, None, &mut one, 1);
for threads in [2, 5, 16] {
let mut y = vec![0f32; m * n];
linear(&x, m, k, &w, n, None, &mut y, threads);
assert!(y.iter().zip(&one).all(|(a, b)| a.to_bits() == b.to_bits()));
}
for i in [0, 36] {
for j in [0, 50, 100] {
let d = dot(&x[i * k..(i + 1) * k], &w[j * k..(j + 1) * k]);
assert_eq!(d.to_bits(), one[i * n + j].to_bits());
}
}
}
}