use rayon::prelude::*;
const MR: usize = 8;
const NR: usize = 8;
pub struct PackedWeight {
p: Vec<f32>,
n: usize,
k: usize,
}
impl PackedWeight {
pub fn new(w: &[f32], n: usize, k: usize) -> Self {
assert_eq!(w.len(), n * k, "weight is not [n, k]");
let panels = n.div_ceil(NR);
let mut p = vec![0f32; panels * k * NR];
for panel in 0..panels {
let base = panel * k * NR;
for j in 0..NR {
let row = panel * NR + j;
if row >= n {
break; }
let src = &w[row * k..(row + 1) * k];
for (kk, &v) in src.iter().enumerate() {
p[base + kk * NR + j] = v;
}
}
}
Self { p, n, k }
}
pub fn n(&self) -> usize {
self.n
}
pub fn k(&self) -> usize {
self.k
}
pub fn unpack(&self) -> Vec<f32> {
let mut w = vec![0f32; self.n * self.k];
for panel in 0..self.n.div_ceil(NR) {
let base = panel * self.k * NR;
for j in 0..NR {
let row = panel * NR + j;
if row >= self.n {
break;
}
for kk in 0..self.k {
w[row * self.k + kk] = self.p[base + kk * NR + j];
}
}
}
w
}
}
pub fn gemv_packed(c: &mut [f32], a: &[f32], w: &PackedWeight, bias: Option<&[f32]>) {
let (n, k) = (w.n, w.k);
assert_eq!(a.len(), k, "lhs is not [1, k]");
assert!(c.len() >= n, "output too small");
if let Some(b) = bias {
assert_eq!(b.len(), n, "bias width");
}
const SERIAL_MACS: usize = 1 << 18;
const TASK_MACS: usize = 1 << 16;
let panels = n.div_ceil(NR);
let per_task = if n * k <= SERIAL_MACS {
panels
} else {
let by_work = TASK_MACS.div_ceil(NR * k.max(1));
let by_threads = panels.div_ceil(4 * rayon::current_num_threads().max(1));
by_work.max(by_threads).max(1).min(panels)
};
c[..n]
.par_chunks_mut(per_task * NR)
.enumerate()
.for_each(|(task, out)| {
for (p, oc) in out.chunks_mut(NR).enumerate() {
let panel = task * per_task + p;
let wp = &w.p[panel * k * NR..(panel + 1) * k * NR];
let mut acc = [0f32; NR];
if let Some(b) = bias {
let j0 = panel * NR;
acc[..oc.len()].copy_from_slice(&b[j0..j0 + oc.len()]);
}
for kk in 0..k {
let av = a[kk];
let wrow = &wp[kk * NR..(kk + 1) * NR];
for j in 0..NR {
acc[j] = av.mul_add(wrow[j], acc[j]);
}
}
oc.copy_from_slice(&acc[..oc.len()]);
}
});
}
pub fn gemm_packed(c: &mut [f32], a: &[f32], w: &PackedWeight, m: usize, bias: Option<&[f32]>) {
let (n, k) = (w.n, w.k);
assert_eq!(a.len(), m * k, "lhs is not [m, k]");
assert!(c.len() >= m * n, "output too small");
if let Some(b) = bias {
assert_eq!(b.len(), n, "bias width");
}
let panels = n.div_ceil(NR);
let use_2d = std::env::var("OSFKB_CPU_GEMM_2D").ok().as_deref() != Some("0");
if use_2d {
let nblocks = m.div_ceil(MR);
let mut apack = vec![0f32; nblocks * k * MR];
apack
.par_chunks_mut(k * MR)
.enumerate()
.for_each(|(blk, ap)| {
let m0 = blk * MR;
let rows = MR.min(m - m0);
for i in 0..rows {
let src = &a[(m0 + i) * k..(m0 + i + 1) * k];
for (kk, &v) in src.iter().enumerate() {
ap[kk * MR + i] = v;
}
}
});
let threads = rayon::current_num_threads().max(1);
let nchunks = (3 * threads).div_ceil(nblocks).clamp(1, panels);
let chunk = panels.div_ceil(nchunks);
let cptr = SendPtr(c.as_mut_ptr());
let panel_major = std::env::var("OSFKB_CPU_GEMM_PM").ok().as_deref() != Some("0");
(0..nblocks * nchunks).into_par_iter().for_each(|task| {
let (blk, ch) = if panel_major {
(task % nblocks, task / nblocks)
} else {
(task / nchunks, task % nchunks)
};
let (p0, p1) = (ch * chunk, ((ch + 1) * chunk).min(panels));
if p0 >= p1 {
return;
}
let m0 = blk * MR;
let rows = MR.min(m - m0);
let ap = &apack[blk * k * MR..(blk + 1) * k * MR];
let mut acc = [0f32; MR * NR];
for panel in p0..p1 {
let wp = &w.p[panel * k * NR..(panel + 1) * k * NR];
let j0 = panel * NR;
let cols = NR.min(n - j0);
let mut bvec = [0f32; NR];
if let Some(b) = bias {
bvec[..cols].copy_from_slice(&b[j0..j0 + cols]);
}
kernel(&mut acc, ap, wp, k, &bvec);
unsafe {
let base = cptr.get();
for i in 0..rows {
std::ptr::copy_nonoverlapping(
acc.as_ptr().add(i * NR),
base.add((m0 + i) * n + j0),
cols,
);
}
}
}
});
return;
}
c[..m * n]
.par_chunks_mut(MR * n)
.enumerate()
.for_each(|(blk, cblk)| {
let m0 = blk * MR;
let rows = cblk.len() / n; let mut ap = vec![0f32; k * MR];
for i in 0..rows {
let src = &a[(m0 + i) * k..(m0 + i + 1) * k];
for (kk, &v) in src.iter().enumerate() {
ap[kk * MR + i] = v;
}
}
let mut acc = [0f32; MR * NR];
for panel in 0..panels {
let wp = &w.p[panel * k * NR..(panel + 1) * k * NR];
let j0 = panel * NR;
let cols = NR.min(n - j0);
let mut bvec = [0f32; NR];
if let Some(b) = bias {
bvec[..cols].copy_from_slice(&b[j0..j0 + cols]);
}
kernel(&mut acc, &ap, wp, k, &bvec);
for i in 0..rows {
cblk[i * n + j0..i * n + j0 + cols]
.copy_from_slice(&acc[i * NR..i * NR + cols]);
}
}
});
}
pub struct PackedWeightI8 {
q: Vec<i8>,
scales: Vec<f32>,
n: usize,
k: usize,
}
impl PackedWeightI8 {
pub fn new(w: &[f32], n: usize, k: usize) -> Self {
assert_eq!(w.len(), n * k, "weight is not [n, k]");
let k8 = k.div_ceil(8);
let panels = n.div_ceil(NR);
let mut scales = vec![0f32; n];
let mut q = vec![0i8; panels * k8 * NR * 8];
for row in 0..n {
let src = &w[row * k..(row + 1) * k];
let absmax = src.iter().fold(0f32, |m, v| m.max(v.abs()));
let s = if absmax > 0.0 { absmax / 127.0 } else { 1.0 };
scales[row] = s;
let inv = 1.0 / s;
let (panel, j) = (row / NR, row % NR);
for (kk, &v) in src.iter().enumerate() {
let c = kk / 8;
let l = kk % 8;
q[((panel * k8 + c) * NR + j) * 8 + l] =
(v * inv).round().clamp(-127.0, 127.0) as i8;
}
}
Self { q, scales, n, k }
}
}
pub fn gemm_i8(c: &mut [f32], a: &[f32], w: &PackedWeightI8, m: usize, bias: Option<&[f32]>) {
let (n, k) = (w.n, w.k);
assert_eq!(a.len(), m * k, "lhs is not [m, k]");
assert!(c.len() >= m * n, "output too small");
let k8 = k.div_ceil(8);
let kb32 = k.div_ceil(32);
let panels = n.div_ceil(NR);
let mpairs = m.div_ceil(2);
let mut qa = vec![0i8; mpairs * k8 * 16];
let mut sa = vec![0f32; m * kb32];
qa.par_chunks_mut(k8 * 16)
.zip(sa.par_chunks_mut(2 * kb32))
.enumerate()
.for_each(|(p, (qp, sp))| {
for r in 0..2usize {
let row = p * 2 + r;
if row >= m {
break;
}
let src = &a[row * k..(row + 1) * k];
for blk in 0..kb32 {
let lo = blk * 32;
let hi = (lo + 32).min(k);
let absmax = src[lo..hi].iter().fold(0f32, |mx, v| mx.max(v.abs()));
let sc = if absmax > 0.0 { absmax / 127.0 } else { 1.0 };
sp[r * kb32 + blk] = sc;
let inv = 1.0 / sc;
for kk in lo..hi {
qp[(kk / 8) * 16 + r * 8 + (kk % 8)] =
(src[kk] * inv).round().clamp(-127.0, 127.0) as i8;
}
}
}
});
let cptr = SendPtr(c.as_mut_ptr());
(0..mpairs * panels).into_par_iter().for_each(|task| {
let mp = task / panels;
let panel = task % panels;
let rows = 2.min(m - mp * 2);
let cols = NR.min(n - panel * NR);
let mut facc = [0f32; 16];
for blk in 0..kb32 {
let c0 = blk * 4;
let c1 = (c0 + 4).min(k8);
let ap = &qa[(mp * k8 + c0) * 16..(mp * k8 + c1) * 16];
let wp = &w.q[(panel * k8 + c0) * NR * 8..(panel * k8 + c1) * NR * 8];
let mut acc = [0i32; 16];
#[cfg(target_arch = "aarch64")]
unsafe {
kernel_smmla(&mut acc, ap, wp, c1 - c0)
};
#[cfg(not(target_arch = "aarch64"))]
kernel_i8_portable(&mut acc, ap, wp, c1 - c0);
for r in 0..rows {
let sblk = sa[(mp * 2 + r) * kb32 + blk];
for j in 0..NR {
facc[r * NR + j] += acc[r * NR + j] as f32 * sblk;
}
}
}
unsafe {
let base = cptr.get();
for r in 0..rows {
let row = mp * 2 + r;
for j in 0..cols {
let col = panel * NR + j;
let mut v = facc[r * NR + j] * w.scales[col];
if let Some(b) = bias {
v += b[col];
}
*base.add(row * n + col) = v;
}
}
}
});
}
pub fn i8mm_available() -> bool {
#[cfg(target_arch = "aarch64")]
{
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
return *V.get_or_init(|| std::arch::is_aarch64_feature_detected!("i8mm"));
}
#[cfg(not(target_arch = "aarch64"))]
false
}
pub struct PackedWeightF16 {
p: Vec<u16>,
n: usize,
k: usize,
}
impl PackedWeightF16 {
pub fn new(w: &[f32], n: usize, k: usize) -> Self {
assert_eq!(w.len(), n * k, "weight is not [n, k]");
let panels = n.div_ceil(NR);
let mut p = vec![0u16; panels * k * NR];
for panel in 0..panels {
let base = panel * k * NR;
for j in 0..NR {
let row = panel * NR + j;
if row >= n {
break; }
let src = &w[row * k..(row + 1) * k];
for (kk, &v) in src.iter().enumerate() {
p[base + kk * NR + j] = f32_to_f16(v);
}
}
}
Self { p, n, k }
}
}
pub fn gemm_f16(c: &mut [f32], a: &[f32], w: &PackedWeightF16, m: usize, bias: Option<&[f32]>) {
let (n, k) = (w.n, w.k);
assert_eq!(a.len(), m * k, "lhs is not [m, k]");
assert!(c.len() >= m * n, "output too small");
if let Some(b) = bias {
assert_eq!(b.len(), n, "bias width");
}
let panels = n.div_ceil(NR);
let nblocks = m.div_ceil(MR);
let mut apack = vec![0u16; nblocks * k * MR];
apack.par_chunks_mut(k * MR).enumerate().for_each_init(
|| vec![0f32; k * MR],
|tmp, (blk, ap)| {
let m0 = blk * MR;
let rows = MR.min(m - m0);
if rows < MR {
tmp.iter_mut().for_each(|v| *v = 0.0); }
for i in 0..rows {
let src = &a[(m0 + i) * k..(m0 + i + 1) * k];
for (kk, &v) in src.iter().enumerate() {
tmp[kk * MR + i] = v;
}
}
f32_to_f16_slice(ap, tmp);
},
);
let threads = rayon::current_num_threads().max(1);
let flat = std::env::var("OSFKB_CPU_F16_FLAT").ok().as_deref() == Some("1");
let nchunks = if flat {
panels
} else {
(3 * threads).div_ceil(nblocks).clamp(1, panels)
};
let chunk = panels.div_ceil(nchunks);
let cptr = SendPtr(c.as_mut_ptr());
(0..nblocks * nchunks).into_par_iter().for_each(|task| {
let (blk, ch) = (task % nblocks, task / nblocks);
let (p0, p1) = (ch * chunk, ((ch + 1) * chunk).min(panels));
if p0 >= p1 {
return;
}
let m0 = blk * MR;
let rows = MR.min(m - m0);
let ap = &apack[blk * k * MR..(blk + 1) * k * MR];
let mut acc = [0f32; MR * NR];
for panel in p0..p1 {
let wp = &w.p[panel * k * NR..(panel + 1) * k * NR];
let j0 = panel * NR;
let cols = NR.min(n - j0);
let mut bvec = [0f32; NR];
if let Some(b) = bias {
bvec[..cols].copy_from_slice(&b[j0..j0 + cols]);
}
#[cfg(target_arch = "aarch64")]
unsafe {
kernel_fmlal(&mut acc, ap, wp, k, &bvec)
};
#[cfg(not(target_arch = "aarch64"))]
kernel_f16_portable(&mut acc, ap, wp, k, &bvec);
unsafe {
let base = cptr.get();
for i in 0..rows {
std::ptr::copy_nonoverlapping(
acc.as_ptr().add(i * NR),
base.add((m0 + i) * n + j0),
cols,
);
}
}
}
});
}
pub fn fhm_available() -> bool {
#[cfg(target_arch = "aarch64")]
{
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
return *V.get_or_init(|| std::arch::is_aarch64_feature_detected!("fhm"));
}
#[cfg(not(target_arch = "aarch64"))]
false
}
pub fn f32_to_f16(x: f32) -> u16 {
let b = x.to_bits();
let sign = ((b >> 16) & 0x8000) as u16;
let exp = (b >> 23) & 0xff;
let man = b & 0x007f_ffff;
if exp == 0xff {
return sign | 0x7c00 | ((man >> 13) as u16) | u16::from(man != 0);
}
let e = exp as i32 - 127 + 15; if e >= 31 {
return sign | 0x7c00; }
if e <= 0 {
if e < -10 {
return sign; }
let sig = man | 0x0080_0000;
let shift = (14 - e) as u32;
let half = 1u32 << (shift - 1);
let rounded = sig + half - 1 + ((sig >> shift) & 1);
return sign | (rounded >> shift) as u16;
}
let mut v = ((e as u32) << 10) | (man >> 13);
let rem = man & 0x1fff;
if rem > 0x1000 || (rem == 0x1000 && (v & 1) == 1) {
v += 1;
}
if v >= 0x7c00 {
return sign | 0x7c00;
}
sign | v as u16
}
pub fn f16_to_f32(bits: u16) -> f32 {
let sign = u32::from(bits >> 15) << 31;
let exp = u32::from(bits >> 10) & 0x1f;
let man = u32::from(bits) & 0x3ff;
if exp == 0x1f {
return f32::from_bits(sign | 0x7f80_0000 | (man << 13));
}
if exp == 0 {
if man == 0 {
return f32::from_bits(sign);
}
let shift = man.leading_zeros() - 21;
let man = (man << shift) & 0x3ff;
let exp = 127 - 15 + 1 - shift;
return f32::from_bits(sign | (exp << 23) | (man << 13));
}
f32::from_bits(sign | ((exp + 127 - 15) << 23) | (man << 13))
}
pub fn f32_to_f16_slice(dst: &mut [u16], src: &[f32]) {
assert_eq!(dst.len(), src.len());
#[cfg_attr(not(target_arch = "aarch64"), allow(unused_mut))]
let mut i = 0;
#[cfg(target_arch = "aarch64")]
{
while i + 8 <= src.len() {
unsafe {
core::arch::asm!(
"ld1 {{v0.4s, v1.4s}}, [{p}]",
"fcvtn v2.4h, v0.4s",
"fcvtn2 v2.8h, v1.4s",
"st1 {{v2.8h}}, [{q}]",
p = in(reg) src.as_ptr().add(i),
q = in(reg) dst.as_mut_ptr().add(i),
out("v0") _, out("v1") _, out("v2") _,
options(nostack)
);
}
i += 8;
}
}
for j in i..src.len() {
dst[j] = f32_to_f16(src[j]);
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "fhm", enable = "fp16")]
unsafe fn kernel_fmlal(
acc: &mut [f32; MR * NR],
ap: &[u16],
wp: &[u16],
k: usize,
bias: &[f32; NR],
) {
use std::arch::aarch64::*;
unsafe {
let blo = vld1q_f32(bias.as_ptr());
let bhi = vld1q_f32(bias.as_ptr().add(4));
let mut c: [float32x4_t; 16] = [blo; 16];
for i in 0..MR {
c[i * 2 + 1] = bhi;
}
let (pa, pw) = (ap.as_ptr(), wp.as_ptr());
for kk in 0..k {
let w = vld1q_u16(pw.add(kk * NR));
let a = vld1q_u16(pa.add(kk * MR));
core::arch::asm!(
"fmlal {c0:v}.4s, {w:v}.4h, {a:v}.h[0]",
"fmlal2 {c1:v}.4s, {w:v}.4h, {a:v}.h[0]",
"fmlal {c2:v}.4s, {w:v}.4h, {a:v}.h[1]",
"fmlal2 {c3:v}.4s, {w:v}.4h, {a:v}.h[1]",
"fmlal {c4:v}.4s, {w:v}.4h, {a:v}.h[2]",
"fmlal2 {c5:v}.4s, {w:v}.4h, {a:v}.h[2]",
"fmlal {c6:v}.4s, {w:v}.4h, {a:v}.h[3]",
"fmlal2 {c7:v}.4s, {w:v}.4h, {a:v}.h[3]",
c0 = inout(vreg) c[0], c1 = inout(vreg) c[1],
c2 = inout(vreg) c[2], c3 = inout(vreg) c[3],
c4 = inout(vreg) c[4], c5 = inout(vreg) c[5],
c6 = inout(vreg) c[6], c7 = inout(vreg) c[7],
w = in(vreg) w, a = in(vreg_low16) a,
options(pure, nomem, nostack)
);
core::arch::asm!(
"fmlal {c0:v}.4s, {w:v}.4h, {a:v}.h[4]",
"fmlal2 {c1:v}.4s, {w:v}.4h, {a:v}.h[4]",
"fmlal {c2:v}.4s, {w:v}.4h, {a:v}.h[5]",
"fmlal2 {c3:v}.4s, {w:v}.4h, {a:v}.h[5]",
"fmlal {c4:v}.4s, {w:v}.4h, {a:v}.h[6]",
"fmlal2 {c5:v}.4s, {w:v}.4h, {a:v}.h[6]",
"fmlal {c6:v}.4s, {w:v}.4h, {a:v}.h[7]",
"fmlal2 {c7:v}.4s, {w:v}.4h, {a:v}.h[7]",
c0 = inout(vreg) c[8], c1 = inout(vreg) c[9],
c2 = inout(vreg) c[10], c3 = inout(vreg) c[11],
c4 = inout(vreg) c[12], c5 = inout(vreg) c[13],
c6 = inout(vreg) c[14], c7 = inout(vreg) c[15],
w = in(vreg) w, a = in(vreg_low16) a,
options(pure, nomem, nostack)
);
}
for i in 0..MR {
vst1q_f32(acc.as_mut_ptr().add(i * NR), c[i * 2]);
vst1q_f32(acc.as_mut_ptr().add(i * NR + 4), c[i * 2 + 1]);
}
}
}
#[cfg(not(target_arch = "aarch64"))]
fn kernel_f16_portable(
acc: &mut [f32; MR * NR],
ap: &[u16],
wp: &[u16],
k: usize,
bias: &[f32; NR],
) {
for i in 0..MR {
for j in 0..NR {
acc[i * NR + j] = bias[j];
}
}
for kk in 0..k {
for i in 0..MR {
let av = f16_to_f32(ap[kk * MR + i]);
for j in 0..NR {
acc[i * NR + j] += av * f16_to_f32(wp[kk * NR + j]);
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "i8mm")]
unsafe fn kernel_smmla(acc: &mut [i32; 16], ap: &[i8], wp: &[i8], k8: usize) {
use std::arch::aarch64::*;
unsafe {
let mut c0 = vdupq_n_s32(0); let mut c1 = vdupq_n_s32(0); let mut c2 = vdupq_n_s32(0); let mut c3 = vdupq_n_s32(0); let (pa, pw) = (ap.as_ptr(), wp.as_ptr());
for c in 0..k8 {
let a = vld1q_s8(pa.add(c * 16)); let w0 = vld1q_s8(pw.add(c * 64)); let w1 = vld1q_s8(pw.add(c * 64 + 16)); let w2 = vld1q_s8(pw.add(c * 64 + 32));
let w3 = vld1q_s8(pw.add(c * 64 + 48));
core::arch::asm!(
"smmla {c0:v}.4s, {a:v}.16b, {w0:v}.16b",
"smmla {c1:v}.4s, {a:v}.16b, {w1:v}.16b",
"smmla {c2:v}.4s, {a:v}.16b, {w2:v}.16b",
"smmla {c3:v}.4s, {a:v}.16b, {w3:v}.16b",
c0 = inout(vreg) c0,
c1 = inout(vreg) c1,
c2 = inout(vreg) c2,
c3 = inout(vreg) c3,
a = in(vreg) a,
w0 = in(vreg) w0,
w1 = in(vreg) w1,
w2 = in(vreg) w2,
w3 = in(vreg) w3,
options(pure, nomem, nostack)
);
}
let mut lanes = [0i32; 4];
for (bi, cc) in [c0, c1, c2, c3].into_iter().enumerate() {
vst1q_s32(lanes.as_mut_ptr(), cc);
acc[bi * 2] = lanes[0];
acc[bi * 2 + 1] = lanes[1];
acc[NR + bi * 2] = lanes[2];
acc[NR + bi * 2 + 1] = lanes[3];
}
}
}
#[cfg_attr(target_arch = "aarch64", allow(dead_code))]
fn kernel_i8_portable(acc: &mut [i32; 16], ap: &[i8], wp: &[i8], k8: usize) {
for c in 0..k8 {
for r in 0..2usize {
for j in 0..NR {
let mut s = 0i32;
for l in 0..8 {
s += ap[c * 16 + r * 8 + l] as i32 * wp[(c * NR + j) * 8 + l] as i32;
}
acc[r * NR + j] += s;
}
}
}
}
struct SendPtr(*mut f32);
impl SendPtr {
fn get(&self) -> *mut f32 {
self.0
}
}
unsafe impl Send for SendPtr {}
unsafe impl Sync for SendPtr {}
#[inline]
fn kernel(acc: &mut [f32; MR * NR], ap: &[f32], wp: &[f32], k: usize, bias: &[f32; NR]) {
#[cfg(target_arch = "aarch64")]
{
unsafe { kernel_neon(acc, ap, wp, k, bias) }
}
#[cfg(target_arch = "x86_64")]
{
use std::sync::OnceLock;
static AVX2: OnceLock<bool> = OnceLock::new();
if *AVX2.get_or_init(|| {
std::arch::is_x86_feature_detected!("avx2")
&& std::arch::is_x86_feature_detected!("fma")
}) {
unsafe { kernel_avx2(acc, ap, wp, k, bias) }
} else {
kernel_portable(acc, ap, wp, k, bias);
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
kernel_portable(acc, ap, wp, k, bias);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn kernel_avx2(
acc: &mut [f32; MR * NR],
ap: &[f32],
wp: &[f32],
k: usize,
bias: &[f32; NR],
) {
use std::arch::x86_64::*;
unsafe {
let b = _mm256_loadu_ps(bias.as_ptr());
let mut c = [b; MR];
let (pa, pw) = (ap.as_ptr(), wp.as_ptr());
for kk in 0..k {
let w = _mm256_loadu_ps(pw.add(kk * NR));
let a = pa.add(kk * MR);
c[0] = _mm256_fmadd_ps(_mm256_set1_ps(*a), w, c[0]);
c[1] = _mm256_fmadd_ps(_mm256_set1_ps(*a.add(1)), w, c[1]);
c[2] = _mm256_fmadd_ps(_mm256_set1_ps(*a.add(2)), w, c[2]);
c[3] = _mm256_fmadd_ps(_mm256_set1_ps(*a.add(3)), w, c[3]);
c[4] = _mm256_fmadd_ps(_mm256_set1_ps(*a.add(4)), w, c[4]);
c[5] = _mm256_fmadd_ps(_mm256_set1_ps(*a.add(5)), w, c[5]);
c[6] = _mm256_fmadd_ps(_mm256_set1_ps(*a.add(6)), w, c[6]);
c[7] = _mm256_fmadd_ps(_mm256_set1_ps(*a.add(7)), w, c[7]);
}
for i in 0..MR {
_mm256_storeu_ps(acc.as_mut_ptr().add(i * NR), c[i]);
}
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn kernel_neon(
acc: &mut [f32; MR * NR],
ap: &[f32],
wp: &[f32],
k: usize,
bias: &[f32; NR],
) {
use std::arch::aarch64::*;
unsafe {
let (b0, b1) = (vld1q_f32(bias.as_ptr()), vld1q_f32(bias.as_ptr().add(4)));
let mut c = [
b0, b1, b0, b1, b0, b1, b0, b1, b0, b1, b0, b1, b0, b1, b0, b1,
];
let (pa, pw) = (ap.as_ptr(), wp.as_ptr());
for kk in 0..k {
let w0 = vld1q_f32(pw.add(kk * NR));
let w1 = vld1q_f32(pw.add(kk * NR + 4));
let a0 = vld1q_f32(pa.add(kk * MR)); let a1 = vld1q_f32(pa.add(kk * MR + 4));
c[0] = vfmaq_laneq_f32(c[0], w0, a0, 0);
c[1] = vfmaq_laneq_f32(c[1], w1, a0, 0);
c[2] = vfmaq_laneq_f32(c[2], w0, a0, 1);
c[3] = vfmaq_laneq_f32(c[3], w1, a0, 1);
c[4] = vfmaq_laneq_f32(c[4], w0, a0, 2);
c[5] = vfmaq_laneq_f32(c[5], w1, a0, 2);
c[6] = vfmaq_laneq_f32(c[6], w0, a0, 3);
c[7] = vfmaq_laneq_f32(c[7], w1, a0, 3);
c[8] = vfmaq_laneq_f32(c[8], w0, a1, 0);
c[9] = vfmaq_laneq_f32(c[9], w1, a1, 0);
c[10] = vfmaq_laneq_f32(c[10], w0, a1, 1);
c[11] = vfmaq_laneq_f32(c[11], w1, a1, 1);
c[12] = vfmaq_laneq_f32(c[12], w0, a1, 2);
c[13] = vfmaq_laneq_f32(c[13], w1, a1, 2);
c[14] = vfmaq_laneq_f32(c[14], w0, a1, 3);
c[15] = vfmaq_laneq_f32(c[15], w1, a1, 3);
}
for i in 0..MR {
vst1q_f32(acc.as_mut_ptr().add(i * NR), c[i * 2]);
vst1q_f32(acc.as_mut_ptr().add(i * NR + 4), c[i * 2 + 1]);
}
}
}
#[cfg_attr(target_arch = "aarch64", allow(dead_code))]
fn kernel_portable(acc: &mut [f32; MR * NR], ap: &[f32], wp: &[f32], k: usize, bias: &[f32; NR]) {
for i in 0..MR {
acc[i * NR..(i + 1) * NR].copy_from_slice(bias);
}
for kk in 0..k {
for i in 0..MR {
let av = ap[kk * MR + i];
for j in 0..NR {
acc[i * NR + j] += av * wp[kk * NR + j];
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn naive(a: &[f32], w: &[f32], m: usize, n: usize, k: usize, bias: Option<&[f32]>) -> Vec<f32> {
let mut c = vec![0f32; m * n];
for i in 0..m {
for j in 0..n {
let mut s = bias.map_or(0.0, |b| b[j]);
for kk in 0..k {
s += a[i * k + kk] * w[j * k + kk];
}
c[i * n + j] = s;
}
}
c
}
#[test]
fn gemv_is_bit_identical_to_the_m1_gemm() {
for (n, k) in [
(1usize, 1usize),
(5, 7),
(8, 16),
(17, 33),
(512, 512), (3072, 1024), (4096, 1024),
] {
let a: Vec<f32> = (0..k)
.map(|i| ((i * 37 % 101) as f32 - 50.0) / 50.0)
.collect();
let w: Vec<f32> = (0..n * k)
.map(|i| ((i * 61 % 197) as f32 - 98.0) / 98.0)
.collect();
let bias: Vec<f32> = (0..n).map(|i| (i % 13) as f32 / 13.0).collect();
let packed = PackedWeight::new(&w, n, k);
for b in [None, Some(bias.as_slice())] {
let mut want = vec![0f32; n];
gemm_packed(&mut want, &a, &packed, 1, b);
let mut got = vec![0f32; n];
gemv_packed(&mut got, &a, &packed, b);
assert_eq!(
want.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"gemv diverged from the m=1 gemm at n={n} k={k} bias={}",
b.is_some()
);
}
}
}
#[test]
fn matches_a_naive_gemm_across_ragged_shapes() {
for (m, n, k) in [
(1usize, 1usize, 1usize),
(3, 5, 7),
(8, 8, 16),
(9, 17, 33),
(64, 768, 768),
(13, 3072, 768),
] {
let a: Vec<f32> = (0..m * k)
.map(|i| ((i % 37) as f32 - 18.0) * 0.031)
.collect();
let w: Vec<f32> = (0..n * k)
.map(|i| ((i % 41) as f32 - 20.0) * 0.017)
.collect();
let bias: Vec<f32> = (0..n).map(|i| (i % 7) as f32 * 0.05).collect();
for b in [None, Some(&bias[..])] {
let want = naive(&a, &w, m, n, k, b);
let packed = PackedWeight::new(&w, n, k);
let mut got = vec![f32::NAN; m * n];
gemm_packed(&mut got, &a, &packed, m, b);
let dmax = got
.iter()
.zip(&want)
.map(|(x, y)| (x - y).abs())
.fold(0f32, f32::max);
assert!(
dmax < 1e-4,
"m={m} n={n} k={k} bias={}: max |Δ| {dmax:.3e}",
b.is_some()
);
}
}
}
#[test]
fn is_deterministic() {
let (m, n, k) = (37usize, 129usize, 71usize);
let a: Vec<f32> = (0..m * k)
.map(|i| ((i % 37) as f32 - 18.0) * 0.031)
.collect();
let w: Vec<f32> = (0..n * k)
.map(|i| ((i % 41) as f32 - 20.0) * 0.017)
.collect();
let packed = PackedWeight::new(&w, n, k);
let run = || {
let mut c = vec![0f32; m * n];
gemm_packed(&mut c, &a, &packed, m, None);
c
};
assert_eq!(
run(),
run(),
"the CPU path is the GPU's oracle: it must not move"
);
}
}
#[cfg(test)]
mod bench {
use super::*;
#[test]
#[ignore = "perf probe: cargo test --release cpu_gemm::bench -- --ignored --nocapture"]
fn beats_the_gemm_crate() {
for (k, n, tag) in [
(768usize, 768usize, "qkv/o "),
(768, 3072, "mlp-up"),
(3072, 768, "mlp-dn"),
] {
let w: Vec<f32> = (0..k * n).map(|i| ((i % 19) as f32 - 9.0) * 0.01).collect();
let packed = PackedWeight::new(&w, n, k);
let bias: Vec<f32> = (0..n).map(|i| (i % 7) as f32 * 0.05).collect();
eprintln!("\n [{tag}] K={k} N={n}");
for m in [16usize, 32, 64, 128, 256, 512] {
let a: Vec<f32> = (0..m * k)
.map(|i| ((i % 23) as f32 - 11.0) * 0.01)
.collect();
let mut y = vec![0f32; m * n];
let (mut ours, mut theirs) = (f64::MAX, f64::MAX);
for _ in 0..15 {
let t0 = std::time::Instant::now();
gemm_packed(&mut y, &a, &packed, m, Some(&bias));
ours = ours.min(t0.elapsed().as_secs_f64());
let t0 = std::time::Instant::now();
unsafe {
gemm::gemm(
m,
n,
k,
y.as_mut_ptr(),
1,
n as isize,
false,
a.as_ptr(),
1,
k as isize,
w.as_ptr(),
k as isize,
1,
0.0,
1.0,
false,
false,
false,
gemm::Parallelism::Rayon(rayon::current_num_threads()),
);
}
theirs = theirs.min(t0.elapsed().as_secs_f64());
}
let gf = |t: f64| 2.0 * (m * n * k) as f64 / t / 1e9;
eprintln!(
" M={m:4} packed {:7.1} GF/s | gemm {:7.1} GF/s -> {:.2}x",
gf(ours),
gf(theirs),
theirs / ours,
);
}
}
}
}