#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
use rayon::prelude::*;
#[cfg(target_arch = "x86_64")]
#[inline]
pub(crate) fn native_available() -> bool {
std::arch::is_x86_feature_detected!("avx512bf16")
&& std::arch::is_x86_feature_detected!("avx512bw")
&& std::arch::is_x86_feature_detected!("avx512f")
}
#[cfg(not(target_arch = "x86_64"))]
#[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
#[inline]
pub(crate) fn native_available() -> bool {
false
}
#[cfg(target_arch = "x86_64")]
const MR: usize = 4;
#[cfg(target_arch = "x86_64")]
const NR: usize = 4;
#[cfg(target_arch = "x86_64")]
const MC: usize = 64;
#[cfg(target_arch = "x86_64")]
pub(crate) fn gemm(a: &[u16], b: &[u16], c: &mut [f32], m: usize, k: usize, n: usize) {
debug_assert_eq!(a.len(), m * k);
debug_assert_eq!(b.len(), k * n);
debug_assert_eq!(c.len(), m * n);
if m == 0 || n == 0 {
return;
}
if k == 0 {
for v in c.iter_mut() {
*v = 0.0;
}
return;
}
let b_t = transpose_b(b, k, n);
c.par_chunks_mut(MC * n)
.enumerate()
.for_each(|(blk, c_blk)| {
let i0 = blk * MC;
let rows = c_blk.len() / n;
let a_blk = &a[i0 * k..i0 * k + rows * k];
unsafe { gemm_block(a_blk, &b_t, c_blk, rows, k, n) };
});
}
#[cfg(target_arch = "x86_64")]
fn transpose_b(b: &[u16], k: usize, n: usize) -> Vec<u16> {
let mut b_t = vec![0u16; n * k];
b_t.par_chunks_mut(k).enumerate().for_each(|(j, dst)| {
for (p, d) in dst.iter_mut().enumerate() {
*d = b[p * n + j];
}
});
b_t
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512bf16,avx512bw,avx512f")]
unsafe fn gemm_block(a: &[u16], b_t: &[u16], c: &mut [f32], rows: usize, k: usize, n: usize) {
let mut i = 0;
while i < rows {
let mr = MR.min(rows - i);
let mut j = 0;
while j < n {
let nr = NR.min(n - j);
unsafe { micro_kernel(a, b_t, c, k, n, i, j, mr, nr) };
j += NR;
}
i += MR;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512bf16,avx512bw,avx512f")]
#[allow(clippy::too_many_arguments)]
#[allow(clippy::needless_range_loop)]
unsafe fn micro_kernel(
a: &[u16],
b_t: &[u16],
c: &mut [f32],
k: usize,
n: usize,
i: usize,
j: usize,
mr: usize,
nr: usize,
) {
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
unsafe {
let mut acc = [[_mm512_setzero_ps(); NR]; MR];
let mut a_ch = [_mm512_setzero_si512(); MR];
let mut b_ch = [_mm512_setzero_si512(); NR];
let a_ptr = a.as_ptr();
let b_ptr = b_t.as_ptr();
let mut p = 0;
while p < k {
let chunk = 32.min(k - p);
if chunk == 32 {
for ii in 0..mr {
a_ch[ii] = _mm512_loadu_si512(a_ptr.add((i + ii) * k + p) as *const __m512i);
}
for jj in 0..nr {
b_ch[jj] = _mm512_loadu_si512(b_ptr.add((j + jj) * k + p) as *const __m512i);
}
} else {
let mask: __mmask32 = (1u32 << chunk) - 1;
for ii in 0..mr {
a_ch[ii] =
_mm512_maskz_loadu_epi16(mask, a_ptr.add((i + ii) * k + p) as *const i16);
}
for jj in 0..nr {
b_ch[jj] =
_mm512_maskz_loadu_epi16(mask, b_ptr.add((j + jj) * k + p) as *const i16);
}
}
for ii in 0..mr {
let av: __m512bh = core::mem::transmute::<__m512i, __m512bh>(a_ch[ii]);
for jj in 0..nr {
let bv: __m512bh = core::mem::transmute::<__m512i, __m512bh>(b_ch[jj]);
acc[ii][jj] = _mm512_dpbf16_ps(acc[ii][jj], av, bv);
}
}
p += 32;
}
for ii in 0..mr {
for jj in 0..nr {
*c.get_unchecked_mut((i + ii) * n + (j + jj)) = _mm512_reduce_add_ps(acc[ii][jj]);
}
}
}
}