use kime_tensor::Epilogue;
use crate::ops::gelu;
use crate::par::{self, Shared};
const LANES: usize = 8;
const BLOCK: usize = 64;
const MR: usize = 6;
pub const NR: usize = 16;
#[cfg(not(target_os = "macos"))]
const NB: usize = 4 * NR;
#[cfg(not(target_os = "macos"))]
const MB: usize = 24 * MR;
#[must_use]
pub fn dot(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len());
#[cfg(target_arch = "aarch64")]
return dot_v::<neon::Neon>(a, b);
#[cfg(target_arch = "x86_64")]
if has_fma() {
return unsafe { dot_fma(a, b) };
}
#[allow(unreachable_code)]
dot_v::<[f32; LANES]>(a, b)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
fn dot_fma(a: &[f32], b: &[f32]) -> f32 {
dot_v::<avx::Avx>(a, b)
}
#[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 splat(v: f32) -> 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 splat(v: f32) -> Self {
[v; 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 splat(v: f32) -> Self {
unsafe { Self(vdupq_n_f32(v), vdupq_n_f32(v)) }
}
#[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_set1_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 splat(v: f32) -> Self {
unsafe { Self(_mm256_set1_ps(v)) }
}
#[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 dot_v<V: V8>(a: &[f32], b: &[f32]) -> f32 {
let k = a.len();
let body = k - k % LANES;
let mut wide = [0f64; LANES];
let mut p = 0;
while p < body {
let end = (p + BLOCK).min(body);
let mut acc = V::zero();
while p < end {
acc = acc.fma(
V::load(a[p..p + LANES].try_into().unwrap()),
V::load(b[p..p + LANES].try_into().unwrap()),
);
p += LANES;
}
for (w, l) in wide.iter_mut().zip(acc.lanes()) {
*w += f64::from(l);
}
}
let v = wide;
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(a[q]).mul_add(f64::from(b[q]), s);
}
s as f32
}
#[must_use]
pub fn pack(w: &[f32], n: usize, k: usize) -> Vec<f32> {
assert_eq!(w.len(), n * k, "w is not [n, k]");
if cfg!(target_os = "macos") { w.to_vec() } else { pack_panels(w, n, k) }
}
fn packed_len(n: usize, k: usize) -> usize {
if cfg!(target_os = "macos") { n * k } else { n.div_ceil(NR) * NR * k }
}
#[must_use]
pub fn scratch_len(k: usize, n: usize) -> usize {
#[cfg(target_os = "macos")]
return blas::ROWS * (k + 3 * n.min(blas::COLS)) + 1;
#[cfg(not(target_os = "macos"))]
{
let _ = (k, n);
0
}
}
#[cfg_attr(target_os = "macos", allow(dead_code))]
fn pack_panels(w: &[f32], n: usize, k: usize) -> Vec<f32> {
let panels = n.div_ceil(NR);
let mut out = vec![0f32; panels * k * NR];
if k == 0 {
return out;
}
for (p, panel) in out.chunks_exact_mut(k * NR).enumerate() {
for c in 0..NR.min(n - p * NR) {
let row = &w[(p * NR + c) * k..][..k];
for (q, &v) in row.iter().enumerate() {
panel[q * NR + c] = v;
}
}
}
out
}
#[inline(always)]
#[cfg_attr(target_os = "macos", allow(dead_code))]
unsafe fn kernel<V: V8, const R: usize>(
x: &[f32],
k: usize,
i: usize,
panel: &[f32],
) -> [[f32; NR]; R] {
let xs: [*const f32; R] = std::array::from_fn(|r| x.as_ptr().wrapping_add((i + r) * k));
let pw = panel.as_ptr();
let mut wide = [[0f64; NR]; R];
let mut q = 0;
while q < k {
let end = (q + BLOCK).min(k);
let mut acc = [[V::zero(); 2]; R];
while q < end {
let (w0, w1) = unsafe {
let at = pw.add(q * NR);
(
V::load(&*at.cast::<[f32; LANES]>()),
V::load(&*at.add(LANES).cast::<[f32; LANES]>()),
)
};
for r in 0..R {
let xv = V::splat(unsafe { *xs[r].add(q) });
acc[r][0] = acc[r][0].fma(xv, w0);
acc[r][1] = acc[r][1].fma(xv, w1);
}
q += 1;
}
for r in 0..R {
for h in 0..2 {
for (w, l) in wide[r][h * LANES..][..LANES].iter_mut().zip(acc[r][h].lanes()) {
*w += f64::from(l);
}
}
}
}
wide.map(|row| row.map(|v| v as f32))
}
#[cfg_attr(target_os = "macos", allow(dead_code))]
struct Args<'a> {
x: &'a [f32],
w: &'a [f32],
b: Option<&'a [f32]>,
ep: Epilogue,
k: usize,
n: usize,
y: &'a Shared<'a>,
}
#[cfg_attr(target_os = "macos", allow(dead_code))]
impl Args<'_> {
#[inline(always)]
fn run<V: V8, const R: usize>(&self, i: usize, p: usize) {
let k = self.k;
let panel = &self.w[p * k * NR..][..k * NR];
let out = unsafe { kernel::<V, R>(self.x, k, i, panel) };
let cols = NR.min(self.n - p * NR);
for (r, row) in out.iter().enumerate() {
for (c, &v) in row[..cols].iter().enumerate() {
self.put(i + r, p * NR + 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) };
}
#[cfg(target_os = "macos")]
fn put_row(&self, i: usize, j0: usize, sums: &[f32]) {
let y = unsafe { self.y.slice_mut(i * self.n + j0, sums.len()) };
let biased = |j: usize, v: f32| match self.b {
Some(b) => v + b[j0 + j],
None => v,
};
match (self.b, self.ep) {
(None, Epilogue::None) => y.copy_from_slice(sums),
(_, Epilogue::None) => {
y.iter_mut().zip(sums).enumerate().for_each(|(j, (y, &v))| *y = biased(j, v))
}
(_, Epilogue::Gelu) => {
y.iter_mut().zip(sums).enumerate().for_each(|(j, (y, &v))| *y = gelu(biased(j, v)))
}
(_, Epilogue::Relu) => y
.iter_mut()
.zip(sums)
.enumerate()
.for_each(|(j, (y, &v))| *y = biased(j, v).max(0.0)),
(Some(b), Epilogue::Accumulate) => {
y.iter_mut().zip(sums).zip(&b[j0..]).for_each(|((y, &v), &b)| *y += v + b)
}
(None, Epilogue::Accumulate) => y.iter_mut().zip(sums).for_each(|(y, &v)| *y += v),
}
}
#[inline(always)]
fn block<V: V8>(&self, rows: (usize, usize), panels: (usize, usize)) {
for p in panels.0..panels.1 {
let mut i = rows.0;
while i + MR <= rows.1 {
self.run::<V, MR>(i, p);
i += MR;
}
match rows.1 - i {
0 => {}
1 => self.run::<V, 1>(i, p),
2 => self.run::<V, 2>(i, p),
3 => self.run::<V, 3>(i, p),
4 => self.run::<V, 4>(i, p),
_ => self.run::<V, 5>(i, p),
}
}
}
fn block_dispatch(&self, rows: (usize, usize), panels: (usize, usize)) {
#[cfg(target_arch = "aarch64")]
return self.block::<neon::Neon>(rows, panels);
#[cfg(target_arch = "x86_64")]
if has_fma() {
unsafe { self.block_fma(rows, panels) };
return;
}
#[allow(unreachable_code)]
self.block::<[f32; LANES]>(rows, panels);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
fn block_fma(&self, rows: (usize, usize), panels: (usize, usize)) {
self.block::<avx::Avx>(rows, panels);
}
}
#[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 w = pack(w, n, k);
let g = Gemm { x, m, k, w: &w, n, b, ep: Epilogue::None };
let len = scratch_len(k, n);
g.run(y, threads, |tasks, f| par::for_each(tasks, threads, |t| f(t, &mut vec![0.0; len])));
}
#[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<'_> {
#[cfg(not(target_os = "macos"))]
fn row_block(&self, threads: usize) -> usize {
let nt = self.n.div_ceil(NB);
[MB, 12 * MR, 6 * MR, 3 * MR]
.into_iter()
.find(|&mb| self.m.div_ceil(mb) * nt >= 3 * threads)
.unwrap_or(MR)
}
pub fn run(
&self,
y: &mut [f32],
threads: usize,
spawn: impl FnOnce(usize, &(dyn Fn(usize, &mut [f32]) + 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(), packed_len(n, k), "w is not [n, k] packed");
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;
}
#[cfg(target_os = "macos")]
{
blas::run(&args, m, threads, spawn);
}
#[cfg(not(target_os = "macos"))]
self.run_panels(&args, threads, spawn);
}
#[cfg(not(target_os = "macos"))]
fn run_panels(
&self,
args: &Args<'_>,
threads: usize,
spawn: impl FnOnce(usize, &(dyn Fn(usize, &mut [f32]) + Sync)),
) {
let (m, n) = (self.m, self.n);
let mb = self.row_block(threads);
let mt = m.div_ceil(mb);
let (panels, per) = (n.div_ceil(NR), NB / NR);
let nt = panels.div_ceil(per);
spawn(mt * nt, &|t, _| {
let (bi, bj) = (t % mt, t / mt);
let rows = (bi * mb, ((bi + 1) * mb).min(m));
args.block_dispatch(rows, (bj * per, ((bj + 1) * per).min(panels)));
});
}
}
#[cfg(target_os = "macos")]
mod blas {
use super::Args;
pub(super) const ROWS: usize = 64;
const KB: usize = 128;
pub(super) const COLS: usize = 256;
#[link(name = "Accelerate", kind = "framework")]
unsafe extern "C" {
fn cblas_sgemm(
order: i32,
trans_a: i32,
trans_b: i32,
m: i32,
n: i32,
k: i32,
alpha: f32,
a: *const f32,
lda: i32,
b: *const f32,
ldb: i32,
beta: f32,
c: *mut f32,
ldc: i32,
);
}
fn f64s(buf: &mut [f32], len: usize) -> &mut [f64] {
let (_, mid, _) = unsafe { buf.align_to_mut::<f64>() };
&mut mid[..len]
}
const ROW_MAJOR: i32 = 101;
const NO_TRANS: i32 = 111;
const TRANS: i32 = 112;
pub(super) fn run(
args: &Args<'_>,
m: usize,
threads: usize,
spawn: impl FnOnce(usize, &(dyn Fn(usize, &mut [f32]) + Sync)),
) {
let (k, n) = (args.k, args.n);
let dim = |v: usize| i32::try_from(v).expect("GEMM sizes fit in an i32");
let _ = threads;
let (mt, cols) = (m.div_ceil(ROWS), COLS.min(n));
spawn(mt * n.div_ceil(cols), &|t, scratch| {
let (bi, bj) = (t % mt, t / mt);
let (r0, rows) = (bi * ROWS, ROWS.min(m - bi * ROWS));
let (j0, nc) = (bj * cols, cols.min(n - bj * cols));
let (xs, rest) = scratch[..ROWS * (k + 3 * cols) + 1].split_at_mut(ROWS * k);
let (c, wide) = rest.split_at_mut(ROWS * cols);
let (c, wide) = (&mut c[..ROWS * nc], &mut f64s(wide, ROWS * cols)[..ROWS * nc]);
wide.fill(0.0);
let x = if rows == ROWS {
&args.x[r0 * k..(r0 + ROWS) * k]
} else {
xs[..rows * k].copy_from_slice(&args.x[r0 * k..(r0 + rows) * k]);
xs[rows * k..].fill(0.0);
&xs[..]
};
let mut p = 0;
while p < k {
let kc = KB.min(k - p);
unsafe {
cblas_sgemm(
ROW_MAJOR,
NO_TRANS,
TRANS,
dim(ROWS),
dim(nc),
dim(kc),
1.0,
x.as_ptr().add(p),
dim(k),
args.w.as_ptr().add(j0 * k + p),
dim(k),
0.0,
c.as_mut_ptr(),
dim(nc),
);
}
wide.iter_mut().zip(c.iter()).for_each(|(w, &v)| *w += f64::from(v));
p += kc;
}
c.iter_mut().zip(wide.iter()).for_each(|(c, &w)| *c = w as f32);
for r in 0..rows {
args.put_row(r0 + r, j0, &c[r * nc..(r + 1) * nc]);
}
});
}
}
#[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);
let tol = if cfg!(target_os = "macos") { 2e-5 } else { 1e-5 };
close(&y, &want, tol, &format!("{m}x{k}x{n}"));
}
}
}
}
#[test]
fn epilogues_on_every_column() {
let mut rng = Rng(5);
let (m, k, n) = (70, 40, 600);
let (x, w, b, y0) = (rng.vec(m * k), rng.vec(n * k), rng.vec(n), rng.vec(m * n));
let packed = pack(&w, n, k);
let lin = naive(&x, m, k, &w, n, Some(&b));
for ep in [Epilogue::None, Epilogue::Gelu, Epilogue::Relu, Epilogue::Accumulate] {
let want: Vec<f32> = lin
.iter()
.zip(&y0)
.map(|(&v, &y)| match ep {
Epilogue::None => v,
Epilogue::Gelu => gelu(v),
Epilogue::Relu => v.max(0.0),
Epilogue::Accumulate => y + v,
})
.collect();
let mut y = y0.clone();
let g = Gemm { x: &x, m, k, w: &packed, n, b: Some(&b), ep };
let len = scratch_len(k, n);
g.run(&mut y, 4, |tasks, f| par::for_each(tasks, 4, |t| f(t, &mut vec![0.0; len])));
close(&y, &want, 1e-4, &format!("{ep:?}"));
}
}
#[test]
fn same_bits_for_any_split() {
let mut rng = Rng(11);
let (m, k, n) = (137, 300, 600);
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, 10, 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, 5, 70, 136] {
let mut row = vec![0f32; n];
linear(&x[i * k..(i + 1) * k], 1, k, &w, n, None, &mut row, 1);
assert!(row.iter().zip(&one[i * n..]).all(|(a, b)| a.to_bits() == b.to_bits()));
}
}
#[test]
fn dot_matches_naive() {
let mut rng = Rng(3);
for k in [0, 1, 7, 8, 64, 65, 200] {
let (a, b) = (rng.vec(k), rng.vec(k));
let want = naive(&a, 1, k, &b, 1, None);
close(&[dot(&a, &b)], &want, 1e-5, &format!("dot {k}"));
}
}
}