1use kime_tensor::Epilogue;
23
24use crate::ops::gelu;
25use crate::par::{self, Shared};
26
27const LANES: usize = 8;
28const BLOCK: usize = 64;
30const MR: usize = 6;
32pub const NR: usize = 16;
34#[cfg(not(target_os = "macos"))]
36const NB: usize = 4 * NR;
37#[cfg(not(target_os = "macos"))]
39const MB: usize = 24 * MR;
40
41#[must_use]
47pub fn dot(a: &[f32], b: &[f32]) -> f32 {
48 assert_eq!(a.len(), b.len());
49 #[cfg(target_arch = "aarch64")]
50 return dot_v::<neon::Neon>(a, b);
51 #[cfg(target_arch = "x86_64")]
52 if has_fma() {
53 return unsafe { dot_fma(a, b) };
55 }
56 #[allow(unreachable_code)]
57 dot_v::<[f32; LANES]>(a, b)
58}
59
60#[cfg(target_arch = "x86_64")]
61#[target_feature(enable = "avx2,fma")]
62fn dot_fma(a: &[f32], b: &[f32]) -> f32 {
63 dot_v::<avx::Avx>(a, b)
64}
65
66#[cfg(target_arch = "x86_64")]
67#[inline]
68fn has_fma() -> bool {
69 std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma")
70}
71
72trait V8: Copy {
75 fn zero() -> Self;
76 fn splat(v: f32) -> Self;
77 fn load(s: &[f32; LANES]) -> Self;
78 fn fma(self, a: Self, b: Self) -> Self;
80 fn lanes(self) -> [f32; LANES];
81}
82
83impl V8 for [f32; LANES] {
84 #[inline(always)]
85 fn zero() -> Self {
86 [0.0; LANES]
87 }
88 #[inline(always)]
89 fn splat(v: f32) -> Self {
90 [v; LANES]
91 }
92 #[inline(always)]
93 fn load(s: &[f32; LANES]) -> Self {
94 *s
95 }
96 #[inline(always)]
97 fn fma(self, a: Self, b: Self) -> Self {
98 std::array::from_fn(|l| a[l].mul_add(b[l], self[l]))
99 }
100 #[inline(always)]
101 fn lanes(self) -> [f32; LANES] {
102 self
103 }
104}
105
106#[cfg(target_arch = "aarch64")]
107mod neon {
108 use std::arch::aarch64::{float32x4_t, vdupq_n_f32, vfmaq_f32, vld1q_f32, vst1q_f32};
109
110 use super::{LANES, V8};
111
112 #[derive(Clone, Copy)]
113 pub(super) struct Neon(float32x4_t, float32x4_t);
114
115 impl V8 for Neon {
116 #[inline(always)]
117 fn zero() -> Self {
118 unsafe { Self(vdupq_n_f32(0.0), vdupq_n_f32(0.0)) }
120 }
121 #[inline(always)]
122 fn splat(v: f32) -> Self {
123 unsafe { Self(vdupq_n_f32(v), vdupq_n_f32(v)) }
125 }
126 #[inline(always)]
127 fn load(s: &[f32; LANES]) -> Self {
128 unsafe { Self(vld1q_f32(s.as_ptr()), vld1q_f32(s.as_ptr().add(4))) }
130 }
131 #[inline(always)]
132 fn fma(self, a: Self, b: Self) -> Self {
133 unsafe { Self(vfmaq_f32(self.0, a.0, b.0), vfmaq_f32(self.1, a.1, b.1)) }
135 }
136 #[inline(always)]
137 fn lanes(self) -> [f32; LANES] {
138 let mut out = [0f32; LANES];
139 unsafe {
141 vst1q_f32(out.as_mut_ptr(), self.0);
142 vst1q_f32(out.as_mut_ptr().add(4), self.1);
143 }
144 out
145 }
146 }
147}
148
149#[cfg(target_arch = "x86_64")]
150mod avx {
151 use std::arch::x86_64::{
152 __m256, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_set1_ps, _mm256_setzero_ps,
153 _mm256_storeu_ps,
154 };
155
156 use super::{LANES, V8};
157
158 #[derive(Clone, Copy)]
160 pub(super) struct Avx(__m256);
161
162 impl V8 for Avx {
163 #[inline(always)]
164 fn zero() -> Self {
165 unsafe { Self(_mm256_setzero_ps()) }
167 }
168 #[inline(always)]
169 fn splat(v: f32) -> Self {
170 unsafe { Self(_mm256_set1_ps(v)) }
172 }
173 #[inline(always)]
174 fn load(s: &[f32; LANES]) -> Self {
175 unsafe { Self(_mm256_loadu_ps(s.as_ptr())) }
177 }
178 #[inline(always)]
179 fn fma(self, a: Self, b: Self) -> Self {
180 unsafe { Self(_mm256_fmadd_ps(a.0, b.0, self.0)) }
182 }
183 #[inline(always)]
184 fn lanes(self) -> [f32; LANES] {
185 let mut out = [0f32; LANES];
186 unsafe { _mm256_storeu_ps(out.as_mut_ptr(), self.0) };
188 out
189 }
190 }
191}
192
193#[inline(always)]
197fn dot_v<V: V8>(a: &[f32], b: &[f32]) -> f32 {
198 let k = a.len();
199 let body = k - k % LANES;
200 let mut wide = [0f64; LANES];
201 let mut p = 0;
202 while p < body {
203 let end = (p + BLOCK).min(body);
204 let mut acc = V::zero();
205 while p < end {
206 acc = acc.fma(
207 V::load(a[p..p + LANES].try_into().unwrap()),
208 V::load(b[p..p + LANES].try_into().unwrap()),
209 );
210 p += LANES;
211 }
212 for (w, l) in wide.iter_mut().zip(acc.lanes()) {
213 *w += f64::from(l);
214 }
215 }
216 let v = wide;
217 let mut s = ((v[0] + v[4]) + (v[2] + v[6])) + ((v[1] + v[5]) + (v[3] + v[7]));
218 for q in body..k {
219 s = f64::from(a[q]).mul_add(f64::from(b[q]), s);
220 }
221 s as f32
222}
223
224#[must_use]
231pub fn pack(w: &[f32], n: usize, k: usize) -> Vec<f32> {
232 assert_eq!(w.len(), n * k, "w is not [n, k]");
233 if cfg!(target_os = "macos") { w.to_vec() } else { pack_panels(w, n, k) }
234}
235
236fn packed_len(n: usize, k: usize) -> usize {
238 if cfg!(target_os = "macos") { n * k } else { n.div_ceil(NR) * NR * k }
239}
240
241#[must_use]
243pub fn scratch_len(k: usize, n: usize) -> usize {
244 #[cfg(target_os = "macos")]
245 return blas::ROWS * (k + 3 * n.min(blas::COLS)) + 1;
246 #[cfg(not(target_os = "macos"))]
247 {
248 let _ = (k, n);
249 0
250 }
251}
252
253#[cfg_attr(target_os = "macos", allow(dead_code))]
255fn pack_panels(w: &[f32], n: usize, k: usize) -> Vec<f32> {
256 let panels = n.div_ceil(NR);
257 let mut out = vec![0f32; panels * k * NR];
258 if k == 0 {
259 return out;
260 }
261 for (p, panel) in out.chunks_exact_mut(k * NR).enumerate() {
262 for c in 0..NR.min(n - p * NR) {
263 let row = &w[(p * NR + c) * k..][..k];
264 for (q, &v) in row.iter().enumerate() {
265 panel[q * NR + c] = v;
266 }
267 }
268 }
269 out
270}
271
272#[inline(always)]
278#[cfg_attr(target_os = "macos", allow(dead_code))]
279unsafe fn kernel<V: V8, const R: usize>(
280 x: &[f32],
281 k: usize,
282 i: usize,
283 panel: &[f32],
284) -> [[f32; NR]; R] {
285 let xs: [*const f32; R] = std::array::from_fn(|r| x.as_ptr().wrapping_add((i + r) * k));
286 let pw = panel.as_ptr();
287 let mut wide = [[0f64; NR]; R];
288 let mut q = 0;
289 while q < k {
290 let end = (q + BLOCK).min(k);
291 let mut acc = [[V::zero(); 2]; R];
292 while q < end {
293 let (w0, w1) = unsafe {
295 let at = pw.add(q * NR);
296 (
297 V::load(&*at.cast::<[f32; LANES]>()),
298 V::load(&*at.add(LANES).cast::<[f32; LANES]>()),
299 )
300 };
301 for r in 0..R {
302 let xv = V::splat(unsafe { *xs[r].add(q) });
304 acc[r][0] = acc[r][0].fma(xv, w0);
305 acc[r][1] = acc[r][1].fma(xv, w1);
306 }
307 q += 1;
308 }
309 for r in 0..R {
310 for h in 0..2 {
311 for (w, l) in wide[r][h * LANES..][..LANES].iter_mut().zip(acc[r][h].lanes()) {
312 *w += f64::from(l);
313 }
314 }
315 }
316 }
317 wide.map(|row| row.map(|v| v as f32))
318}
319
320#[cfg_attr(target_os = "macos", allow(dead_code))]
321struct Args<'a> {
322 x: &'a [f32],
323 w: &'a [f32],
324 b: Option<&'a [f32]>,
325 ep: Epilogue,
326 k: usize,
327 n: usize,
328 y: &'a Shared<'a>,
329}
330
331#[cfg_attr(target_os = "macos", allow(dead_code))]
332impl Args<'_> {
333 #[inline(always)]
334 fn run<V: V8, const R: usize>(&self, i: usize, p: usize) {
335 let k = self.k;
336 let panel = &self.w[p * k * NR..][..k * NR];
337 let out = unsafe { kernel::<V, R>(self.x, k, i, panel) };
339 let cols = NR.min(self.n - p * NR);
340 for (r, row) in out.iter().enumerate() {
341 for (c, &v) in row[..cols].iter().enumerate() {
342 self.put(i + r, p * NR + c, v);
343 }
344 }
345 }
346
347 #[inline(always)]
349 fn put(&self, i: usize, j: usize, v: f32) {
350 let v = match self.b {
351 Some(b) => v + b[j],
352 None => v,
353 };
354 let at = i * self.n + j;
355 let v = match self.ep {
356 Epilogue::None => v,
357 Epilogue::Gelu => gelu(v),
358 Epilogue::Relu => v.max(0.0),
359 Epilogue::Accumulate => v + unsafe { self.y.get(at) },
361 };
362 unsafe { self.y.set(at, v) };
364 }
365
366 #[cfg(target_os = "macos")]
369 fn put_row(&self, i: usize, j0: usize, sums: &[f32]) {
370 let y = unsafe { self.y.slice_mut(i * self.n + j0, sums.len()) };
372 let biased = |j: usize, v: f32| match self.b {
373 Some(b) => v + b[j0 + j],
374 None => v,
375 };
376 match (self.b, self.ep) {
377 (None, Epilogue::None) => y.copy_from_slice(sums),
378 (_, Epilogue::None) => {
379 y.iter_mut().zip(sums).enumerate().for_each(|(j, (y, &v))| *y = biased(j, v))
380 }
381 (_, Epilogue::Gelu) => {
382 y.iter_mut().zip(sums).enumerate().for_each(|(j, (y, &v))| *y = gelu(biased(j, v)))
383 }
384 (_, Epilogue::Relu) => y
385 .iter_mut()
386 .zip(sums)
387 .enumerate()
388 .for_each(|(j, (y, &v))| *y = biased(j, v).max(0.0)),
389 (Some(b), Epilogue::Accumulate) => {
390 y.iter_mut().zip(sums).zip(&b[j0..]).for_each(|((y, &v), &b)| *y += v + b)
391 }
392 (None, Epilogue::Accumulate) => y.iter_mut().zip(sums).for_each(|(y, &v)| *y += v),
393 }
394 }
395
396 #[inline(always)]
398 fn block<V: V8>(&self, rows: (usize, usize), panels: (usize, usize)) {
399 for p in panels.0..panels.1 {
400 let mut i = rows.0;
401 while i + MR <= rows.1 {
402 self.run::<V, MR>(i, p);
403 i += MR;
404 }
405 match rows.1 - i {
406 0 => {}
407 1 => self.run::<V, 1>(i, p),
408 2 => self.run::<V, 2>(i, p),
409 3 => self.run::<V, 3>(i, p),
410 4 => self.run::<V, 4>(i, p),
411 _ => self.run::<V, 5>(i, p),
412 }
413 }
414 }
415
416 fn block_dispatch(&self, rows: (usize, usize), panels: (usize, usize)) {
417 #[cfg(target_arch = "aarch64")]
418 return self.block::<neon::Neon>(rows, panels);
419 #[cfg(target_arch = "x86_64")]
420 if has_fma() {
421 unsafe { self.block_fma(rows, panels) };
423 return;
424 }
425 #[allow(unreachable_code)]
426 self.block::<[f32; LANES]>(rows, panels);
427 }
428
429 #[cfg(target_arch = "x86_64")]
430 #[target_feature(enable = "avx2,fma")]
431 fn block_fma(&self, rows: (usize, usize), panels: (usize, usize)) {
432 self.block::<avx::Avx>(rows, panels);
433 }
434}
435
436#[allow(clippy::too_many_arguments)]
444pub fn linear(
445 x: &[f32],
446 m: usize,
447 k: usize,
448 w: &[f32],
449 n: usize,
450 b: Option<&[f32]>,
451 y: &mut [f32],
452 threads: usize,
453) {
454 let w = pack(w, n, k);
455 let g = Gemm { x, m, k, w: &w, n, b, ep: Epilogue::None };
456 let len = scratch_len(k, n);
457 g.run(y, threads, |tasks, f| par::for_each(tasks, threads, |t| f(t, &mut vec![0.0; len])));
458}
459
460#[derive(Debug, Clone, Copy)]
464pub struct Gemm<'a> {
465 pub x: &'a [f32],
467 pub m: usize,
469 pub k: usize,
471 pub w: &'a [f32],
473 pub n: usize,
475 pub b: Option<&'a [f32]>,
477 pub ep: Epilogue,
479}
480
481impl Gemm<'_> {
482 #[cfg(not(target_os = "macos"))]
484 fn row_block(&self, threads: usize) -> usize {
485 let nt = self.n.div_ceil(NB);
486 [MB, 12 * MR, 6 * MR, 3 * MR]
487 .into_iter()
488 .find(|&mb| self.m.div_ceil(mb) * nt >= 3 * threads)
489 .unwrap_or(MR)
490 }
491
492 pub fn run(
499 &self,
500 y: &mut [f32],
501 threads: usize,
502 spawn: impl FnOnce(usize, &(dyn Fn(usize, &mut [f32]) + Sync)),
503 ) {
504 let Self { x, m, k, w, n, b, ep } = *self;
505 assert_eq!(x.len(), m * k, "x is not [m, k]");
506 assert_eq!(w.len(), packed_len(n, k), "w is not [n, k] packed");
507 assert_eq!(y.len(), m * n, "y is not [m, n]");
508 if let Some(b) = b {
509 assert_eq!(b.len(), n, "b is not [n]");
510 }
511 if m == 0 || n == 0 {
512 return;
513 }
514 let shared = Shared::new(y);
515 let args = Args { x, w, b, ep, k, n, y: &shared };
516 if k == 0 {
517 for i in 0..m {
518 for j in 0..n {
519 args.put(i, j, 0.0);
520 }
521 }
522 return;
523 }
524 #[cfg(target_os = "macos")]
525 {
526 blas::run(&args, m, threads, spawn);
527 }
528 #[cfg(not(target_os = "macos"))]
529 self.run_panels(&args, threads, spawn);
530 }
531
532 #[cfg(not(target_os = "macos"))]
533 fn run_panels(
534 &self,
535 args: &Args<'_>,
536 threads: usize,
537 spawn: impl FnOnce(usize, &(dyn Fn(usize, &mut [f32]) + Sync)),
538 ) {
539 let (m, n) = (self.m, self.n);
540 let mb = self.row_block(threads);
541 let mt = m.div_ceil(mb);
542 let (panels, per) = (n.div_ceil(NR), NB / NR);
543 let nt = panels.div_ceil(per);
544 spawn(mt * nt, &|t, _| {
545 let (bi, bj) = (t % mt, t / mt);
546 let rows = (bi * mb, ((bi + 1) * mb).min(m));
547 args.block_dispatch(rows, (bj * per, ((bj + 1) * per).min(panels)));
548 });
549 }
550}
551
552#[cfg(target_os = "macos")]
553mod blas {
554 use super::Args;
555
556 pub(super) const ROWS: usize = 64;
558 const KB: usize = 128;
561 pub(super) const COLS: usize = 256;
563
564 #[link(name = "Accelerate", kind = "framework")]
565 unsafe extern "C" {
566 fn cblas_sgemm(
567 order: i32,
568 trans_a: i32,
569 trans_b: i32,
570 m: i32,
571 n: i32,
572 k: i32,
573 alpha: f32,
574 a: *const f32,
575 lda: i32,
576 b: *const f32,
577 ldb: i32,
578 beta: f32,
579 c: *mut f32,
580 ldc: i32,
581 );
582 }
583
584 fn f64s(buf: &mut [f32], len: usize) -> &mut [f64] {
586 let (_, mid, _) = unsafe { buf.align_to_mut::<f64>() };
588 &mut mid[..len]
589 }
590
591 const ROW_MAJOR: i32 = 101;
592 const NO_TRANS: i32 = 111;
593 const TRANS: i32 = 112;
594
595 pub(super) fn run(
601 args: &Args<'_>,
602 m: usize,
603 threads: usize,
604 spawn: impl FnOnce(usize, &(dyn Fn(usize, &mut [f32]) + Sync)),
605 ) {
606 let (k, n) = (args.k, args.n);
607 let dim = |v: usize| i32::try_from(v).expect("GEMM sizes fit in an i32");
608 let _ = threads;
609 let (mt, cols) = (m.div_ceil(ROWS), COLS.min(n));
610 spawn(mt * n.div_ceil(cols), &|t, scratch| {
611 let (bi, bj) = (t % mt, t / mt);
612 let (r0, rows) = (bi * ROWS, ROWS.min(m - bi * ROWS));
613 let (j0, nc) = (bj * cols, cols.min(n - bj * cols));
614 let (xs, rest) = scratch[..ROWS * (k + 3 * cols) + 1].split_at_mut(ROWS * k);
615 let (c, wide) = rest.split_at_mut(ROWS * cols);
616 let (c, wide) = (&mut c[..ROWS * nc], &mut f64s(wide, ROWS * cols)[..ROWS * nc]);
617 wide.fill(0.0);
618 let x = if rows == ROWS {
619 &args.x[r0 * k..(r0 + ROWS) * k]
620 } else {
621 xs[..rows * k].copy_from_slice(&args.x[r0 * k..(r0 + rows) * k]);
622 xs[rows * k..].fill(0.0);
623 &xs[..]
624 };
625 let mut p = 0;
626 while p < k {
627 let kc = KB.min(k - p);
628 unsafe {
631 cblas_sgemm(
632 ROW_MAJOR,
633 NO_TRANS,
634 TRANS,
635 dim(ROWS),
636 dim(nc),
637 dim(kc),
638 1.0,
639 x.as_ptr().add(p),
640 dim(k),
641 args.w.as_ptr().add(j0 * k + p),
642 dim(k),
643 0.0,
644 c.as_mut_ptr(),
645 dim(nc),
646 );
647 }
648 wide.iter_mut().zip(c.iter()).for_each(|(w, &v)| *w += f64::from(v));
649 p += kc;
650 }
651 c.iter_mut().zip(wide.iter()).for_each(|(c, &w)| *c = w as f32);
652 for r in 0..rows {
653 args.put_row(r0 + r, j0, &c[r * nc..(r + 1) * nc]);
654 }
655 });
656 }
657}
658
659#[cfg(test)]
660mod tests {
661 use super::*;
662 use crate::testing::{Rng, close};
663
664 fn naive(x: &[f32], m: usize, k: usize, w: &[f32], n: usize, b: Option<&[f32]>) -> Vec<f32> {
665 let mut y = vec![0f32; m * n];
666 for i in 0..m {
667 for j in 0..n {
668 let s: f64 =
669 (0..k).map(|p| f64::from(x[i * k + p]) * f64::from(w[j * k + p])).sum();
670 y[i * n + j] = (s + b.map_or(0.0, |b| f64::from(b[j]))) as f32;
671 }
672 }
673 y
674 }
675
676 #[test]
677 fn matches_naive_on_awkward_shapes() {
678 let mut rng = Rng(7);
679 let shapes = [
680 (0, 8, 5),
681 (1, 1, 1),
682 (1, 7, 3),
683 (3, 16, 2),
684 (4, 9, 3),
685 (5, 64, 7),
686 (130, 33, 50),
687 (129, 1028, 49),
688 (17, 0, 4),
689 ];
690 for (m, k, n) in shapes {
691 let x = rng.vec(m * k);
692 let w = rng.vec(n * k);
693 let b = rng.vec(n);
694 for bias in [None, Some(&b[..])] {
695 let want = naive(&x, m, k, &w, n, bias);
696 for threads in [1, 3] {
697 let mut y = vec![f32::NAN; m * n];
698 linear(&x, m, k, &w, n, bias, &mut y, threads);
699 let tol = if cfg!(target_os = "macos") { 2e-5 } else { 1e-5 };
701 close(&y, &want, tol, &format!("{m}x{k}x{n}"));
702 }
703 }
704 }
705 }
706
707 #[test]
708 fn epilogues_on_every_column() {
709 let mut rng = Rng(5);
710 let (m, k, n) = (70, 40, 600);
711 let (x, w, b, y0) = (rng.vec(m * k), rng.vec(n * k), rng.vec(n), rng.vec(m * n));
712 let packed = pack(&w, n, k);
713 let lin = naive(&x, m, k, &w, n, Some(&b));
714 for ep in [Epilogue::None, Epilogue::Gelu, Epilogue::Relu, Epilogue::Accumulate] {
715 let want: Vec<f32> = lin
716 .iter()
717 .zip(&y0)
718 .map(|(&v, &y)| match ep {
719 Epilogue::None => v,
720 Epilogue::Gelu => gelu(v),
721 Epilogue::Relu => v.max(0.0),
722 Epilogue::Accumulate => y + v,
723 })
724 .collect();
725 let mut y = y0.clone();
726 let g = Gemm { x: &x, m, k, w: &packed, n, b: Some(&b), ep };
727 let len = scratch_len(k, n);
728 g.run(&mut y, 4, |tasks, f| par::for_each(tasks, 4, |t| f(t, &mut vec![0.0; len])));
729 close(&y, &want, 1e-4, &format!("{ep:?}"));
730 }
731 }
732
733 #[test]
734 fn same_bits_for_any_split() {
735 let mut rng = Rng(11);
736 let (m, k, n) = (137, 300, 600);
737 let x = rng.vec(m * k);
738 let w = rng.vec(n * k);
739 let mut one = vec![0f32; m * n];
740 linear(&x, m, k, &w, n, None, &mut one, 1);
741 for threads in [2, 5, 10, 16] {
742 let mut y = vec![0f32; m * n];
743 linear(&x, m, k, &w, n, None, &mut y, threads);
744 assert!(y.iter().zip(&one).all(|(a, b)| a.to_bits() == b.to_bits()));
745 }
746 for i in [0, 5, 70, 136] {
748 let mut row = vec![0f32; n];
749 linear(&x[i * k..(i + 1) * k], 1, k, &w, n, None, &mut row, 1);
750 assert!(row.iter().zip(&one[i * n..]).all(|(a, b)| a.to_bits() == b.to_bits()));
751 }
752 }
753
754 #[test]
755 fn dot_matches_naive() {
756 let mut rng = Rng(3);
757 for k in [0, 1, 7, 8, 64, 65, 200] {
758 let (a, b) = (rng.vec(k), rng.vec(k));
759 let want = naive(&a, 1, k, &b, 1, None);
760 close(&[dot(&a, &b)], &want, 1e-5, &format!("dot {k}"));
761 }
762 }
763}