use crate::{
align::Alignment,
arch::SimdArch,
kernel::SimdKernel,
scalar::Scalar,
vec::AlignedVec,
view::{SimdError, SimdView},
};
pub(super) const GEMM_PACK_B_BYTES_THRESHOLD: usize = 512 * 1024;
#[inline(never)]
fn check_tiled_gemm_dimensions(
a_len: usize,
b_len: usize,
c_len: usize,
m: usize,
n: usize,
k: usize,
) -> Result<(), SimdError> {
let a_needed = super::dims::checked_area(m, k).ok_or(SimdError::LengthMismatch)?;
let b_needed = super::dims::checked_area(k, n).ok_or(SimdError::LengthMismatch)?;
let c_needed = super::dims::checked_area(m, n).ok_or(SimdError::LengthMismatch)?;
if a_len < a_needed || b_len < b_needed || c_len < c_needed {
return Err(SimdError::LengthMismatch);
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn gemm_register_tile<T, Arch, const TILE_M: usize, const TILE_N: usize>(
a_slice: &[T],
c: &mut [T],
r: usize,
current_tile_m: usize,
col_n: usize,
n: usize,
k: usize,
b_base: *const T,
b_row_stride: usize,
) where
Arch: SimdArch + SimdKernel<T>,
T: Scalar,
{
let lane_count = Arch::LANE_COUNT;
let mut accumulators = [[unsafe { Arch::zero() }; TILE_N]; TILE_M];
for i in 0..current_tile_m {
let row_idx = r + i;
for j in 0..TILE_N {
let c_ptr = unsafe { c.as_ptr().add(row_idx * n + col_n + j * lane_count) };
accumulators[i][j] = unsafe { Arch::load_unaligned(c_ptr) };
}
}
for kk in 0..k {
let mut b_regs = [unsafe { Arch::zero() }; TILE_N];
for j in 0..TILE_N {
let b_ptr = unsafe { b_base.add(kk * b_row_stride + j * lane_count) };
b_regs[j] = unsafe { Arch::load_unaligned(b_ptr) };
}
for i in 0..current_tile_m {
let a_reg = unsafe { Arch::splat(a_slice[(r + i) * k + kk]) };
for j in 0..TILE_N {
accumulators[i][j] = unsafe { Arch::fmadd(a_reg, b_regs[j], accumulators[i][j]) };
}
}
}
for i in 0..current_tile_m {
let row_idx = r + i;
for j in 0..TILE_N {
let c_ptr = unsafe { c.as_mut_ptr().add(row_idx * n + col_n + j * lane_count) };
unsafe { Arch::store_unaligned(c_ptr, accumulators[i][j]) };
}
}
}
#[inline]
pub(super) fn gemm_impl<T, Arch, Align, const TILE_M: usize, const TILE_N: usize>(
a: &SimdView<'_, T, Arch, Align>,
b: &SimdView<'_, T, Arch, Align>,
c: &mut [T],
m: usize,
n: usize,
k: usize,
) -> Result<(), SimdError>
where
Arch: SimdArch + SimdKernel<T>,
Align: Alignment,
T: Scalar,
{
struct AssertGEMM<const M: usize, const N: usize>;
impl<const M: usize, const N: usize> AssertGEMM<M, N> {
const OK: () = {
assert!(M >= 1, "TILE_M must be at least 1");
assert!(N >= 1, "TILE_N must be at least 1");
assert!(M * N <= 64, "TILE_M * TILE_N must be <= 64");
};
}
let _ = AssertGEMM::<TILE_M, TILE_N>::OK;
check_tiled_gemm_dimensions(a.len(), b.len(), c.len(), m, n, k)?;
let a_slice = a.as_slice();
let b_slice = b.as_slice();
let lane_count = Arch::LANE_COUNT;
let block_n = TILE_N * lane_count;
let simd_n_len = (n / block_n) * block_n;
let b_bytes = n
.saturating_mul(k)
.saturating_mul(core::mem::size_of::<T>());
if m > TILE_M && simd_n_len > 0 && b_bytes >= GEMM_PACK_B_BYTES_THRESHOLD {
let mut packed = AlignedVec::<T, crate::align::Aligned<64>>::with_capacity(k * block_n);
unsafe {
packed.set_len(k * block_n);
}
let mut col_n = 0;
while col_n < simd_n_len {
for kk in 0..k {
unsafe {
core::ptr::copy_nonoverlapping(
b_slice.as_ptr().add(kk * n + col_n),
packed.as_mut_ptr().add(kk * block_n),
block_n,
);
}
}
let mut r = 0;
while r < m {
let current_tile_m = if r + TILE_M <= m { TILE_M } else { m - r };
unsafe {
gemm_register_tile::<T, Arch, TILE_M, TILE_N>(
a_slice,
c,
r,
current_tile_m,
col_n,
n,
k,
packed.as_ptr(),
block_n,
);
}
r += TILE_M;
}
col_n += block_n;
}
} else {
let mut r = 0;
while r < m {
let current_tile_m = if r + TILE_M <= m { TILE_M } else { m - r };
let mut col_n = 0;
while col_n < simd_n_len {
unsafe {
gemm_register_tile::<T, Arch, TILE_M, TILE_N>(
a_slice,
c,
r,
current_tile_m,
col_n,
n,
k,
b_slice.as_ptr().add(col_n),
n,
);
}
col_n += block_n;
}
r += TILE_M;
}
}
let mut col = simd_n_len;
while col < n {
let w = core::cmp::min(lane_count, n - col);
let mask = unsafe { Arch::leading_k_mask(w) };
let mut r = 0;
while r < m {
let current_tile_m = if r + TILE_M <= m { TILE_M } else { m - r };
let mut acc = [unsafe { Arch::zero() }; TILE_M];
for (i, slot) in acc.iter_mut().take(current_tile_m).enumerate() {
*slot = unsafe {
Arch::masked_load_unaligned(
c.as_ptr().add((r + i) * n + col),
mask,
Arch::zero(),
)
};
}
for kk in 0..k {
let b_reg = unsafe {
Arch::masked_load_unaligned(
b_slice.as_ptr().add(kk * n + col),
mask,
Arch::zero(),
)
};
for (i, slot) in acc.iter_mut().take(current_tile_m).enumerate() {
unsafe {
let a_reg = Arch::splat(a_slice[(r + i) * k + kk]);
*slot = Arch::fmadd(a_reg, b_reg, *slot);
}
}
}
for (i, slot) in acc.iter().take(current_tile_m).enumerate() {
unsafe {
Arch::masked_store_unaligned(c.as_mut_ptr().add((r + i) * n + col), mask, *slot)
};
}
r += TILE_M;
}
col += lane_count;
}
Ok(())
}