use crate::kernel::FloatGemm;
#[cfg(feature = "half")]
use crate::kernel::MixedGemm;
use crate::kernel::epilogue::Epilogue;
use crate::parallel::{self, JobCursor, Parallelism, Ptr};
use crate::scalar::Float;
#[cfg(feature = "half")]
use crate::scalar::NarrowFloat;
#[cfg(any(feature = "half", feature = "int8"))]
use crate::simd::KernelSimd;
use crate::simd::SimdOps;
use crate::workspace::Workspace;
const MT: usize = 4;
const NT: usize = 4;
fn tile_sweep(
m: usize,
n: usize,
bytes: usize,
par: Parallelism,
body: impl Fn(usize, usize) + Copy + Send + Sync,
) {
let n_row_tiles = m.div_ceil(MT);
let n_col_tiles = n.div_ceil(NT);
let n_tiles = n_row_tiles * n_col_tiles;
let n_threads = par.resolve_bandwidth(bytes, n_tiles);
if n_threads <= 1 {
body(0, n_tiles);
return;
}
let cur = JobCursor::new(n_tiles, parallel::job_grain(n_tiles, n_threads));
parallel::for_each_worker(n_threads, |_tid| {
while let Some((s, e)) = cur.next_chunk() {
body(s, e);
}
});
}
#[inline]
fn packed_line_stride<T>(k: usize) -> usize {
let lane = (64 / core::mem::size_of::<T>().max(1)).max(1);
let lines = k.div_ceil(lane).max(1);
let odd_lines = if lines.is_multiple_of(2) {
lines + 1
} else {
lines
};
odd_lines * lane
}
#[inline]
unsafe fn pack_k_contiguous<T: Copy>(
dst: *mut T,
src: *const T,
lead: usize,
k: usize,
dst_stride: usize,
lead_stride: isize,
depth_stride: isize,
) {
unsafe {
let tile = crate::tuning::pack_transpose_tile();
let mut t0 = 0;
while t0 < k {
let te = core::cmp::min(t0 + tile, k);
for t in t0..te {
let col = src.offset(t as isize * depth_stride);
for l in 0..lead {
*dst.add(l * dst_stride + t) = *col.offset(l as isize * lead_stride);
}
}
t0 = te;
}
}
}
#[allow(clippy::too_many_arguments)]
unsafe fn prepack_operands<T: Copy>(
ws: &mut Workspace,
m: usize,
k: usize,
n: usize,
a: *const T,
rsa: isize,
csa: isize,
b: *const T,
rsb: isize,
csb: isize,
) -> (*const T, isize, isize, *const T, isize, isize) {
let pack_a = csa != 1;
let pack_b = rsb != 1;
if !pack_a && !pack_b {
return (a, rsa, csa, b, rsb, csb);
}
let stride = packed_line_stride::<T>(k);
unsafe {
let a_elems = if pack_a { m.saturating_mul(stride) } else { 0 };
let b_elems = if pack_b { n.saturating_mul(stride) } else { 0 };
let r = ws.regions::<T>(a_elems, 1, b_elems);
let (mut a, mut rsa, mut csa) = (a, rsa, csa);
let (mut b, mut rsb, mut csb) = (b, rsb, csb);
if pack_a {
pack_k_contiguous::<T>(r.a_base, a, m, k, stride, rsa, csa);
a = r.a_base;
rsa = stride as isize;
csa = 1;
}
if pack_b {
pack_k_contiguous::<T>(r.b_base, b, n, k, stride, csb, rsb);
b = r.b_base;
rsb = 1;
csb = stride as isize;
}
(a, rsa, csa, b, rsb, csb)
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_epi<T, S, E>(
simd: S,
m: usize,
k: usize,
n: usize,
par: Parallelism,
ws: &mut Workspace,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
b: *const T,
rsb: isize,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
epi: &E,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
E: Epilogue<FloatGemm<T>>,
{
let epi = *epi;
unsafe {
let (a, rsa, csa, b, rsb, csb) =
prepack_operands::<T>(ws, m, k, n, a, rsa, csa, b, rsb, csb);
debug_assert!(
csa == 1 && rsb == 1,
"small_mn kernel requires A rows / B cols unit-stride along k"
);
let n_row_tiles = m.div_ceil(MT);
let sizeof = core::mem::size_of::<T>();
let bytes = m
.saturating_mul(k)
.saturating_add(k.saturating_mul(n))
.saturating_add(m.saturating_mul(n))
.saturating_mul(sizeof);
let a = Ptr(a as *mut T);
let b = Ptr(b as *mut T);
let c = Ptr(c);
let body = move |q_start: usize, q_end: usize| {
let (a, b, c, epi) = (a, b, c, epi);
let a = a.0 as *const T;
let b = b.0 as *const T;
let c = c.0;
simd.vectorize(|| {
for q in q_start..q_end {
let it = q % n_row_tiles;
let jt = q / n_row_tiles;
let i0 = it * MT;
let j0 = jt * NT;
let mi = core::cmp::min(MT, m - i0);
let nj = core::cmp::min(NT, n - j0);
if mi == MT && nj == NT {
full_tile::<T, S, E, MT, NT>(
simd, k, i0, j0, alpha, a, rsa, b, csb, beta, c, rsc, csc, &epi,
);
} else {
for cc in 0..nj {
for ir in 0..mi {
cell_dot::<T, S, E>(
simd,
k,
i0 + ir,
j0 + cc,
alpha,
a,
rsa,
b,
csb,
beta,
c,
rsc,
csc,
&epi,
);
}
}
}
}
});
};
tile_sweep(m, n, bytes, par, body);
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn full_tile<T, S, E, const MT: usize, const NT: usize>(
simd: S,
k: usize,
i0: usize,
j0: usize,
alpha: T,
a: *const T,
rsa: isize,
b: *const T,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
epi: &E,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
E: Epilogue<FloatGemm<T>>,
{
unsafe {
let lanes = <S as SimdOps<T>>::LANES;
let rows: [*const T; MT] = core::array::from_fn(|r| a.offset((i0 + r) as isize * rsa));
let cols: [*const T; NT] = core::array::from_fn(|cc| b.offset((j0 + cc) as isize * csb));
let mut acc = [[simd.zero(); MT]; NT];
let mut kk = 0;
while kk + lanes <= k {
let av: [S::Reg; MT] = core::array::from_fn(|r| simd.loadu(rows[r].add(kk)));
for cc in 0..NT {
let bv = simd.loadu(cols[cc].add(kk));
for r in 0..MT {
acc[cc][r] = simd.mul_add(av[r], bv, acc[cc][r]);
}
}
kk += lanes;
}
for cc in 0..NT {
for r in 0..MT {
let mut dot = simd.reduce_sum(acc[cc][r]);
let mut t = kk;
while t < k {
dot = (*rows[r].add(t)).mul_add(*cols[cc].add(t), dot);
t += 1;
}
let cp = c.offset((i0 + r) as isize * rsc + (j0 + cc) as isize * csc);
let ov = if beta == T::ZERO {
T::ZERO
} else if beta == T::ONE {
*cp
} else {
beta * *cp
};
let out = alpha.mul_add(dot, ov);
*cp = if E::IS_IDENTITY {
out
} else {
epi.apply(out, i0 + r, j0 + cc)
};
}
}
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn cell_dot<T, S, E>(
simd: S,
k: usize,
i: usize,
j: usize,
alpha: T,
a: *const T,
rsa: isize,
b: *const T,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
epi: &E,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
E: Epilogue<FloatGemm<T>>,
{
unsafe {
let row = a.offset(i as isize * rsa); let col = b.offset(j as isize * csb); let dot = super::dot_contiguous::<T, S>(simd, k, row, col);
let cp = c.offset(i as isize * rsc + j as isize * csc);
let ov = if beta == T::ZERO {
T::ZERO
} else if beta == T::ONE {
*cp
} else {
beta * *cp
};
let out = alpha.mul_add(dot, ov);
*cp = if E::IS_IDENTITY {
out
} else {
epi.apply(out, i, j)
};
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_mixed_epi<N, S, E>(
simd: S,
m: usize,
k: usize,
n: usize,
par: Parallelism,
ws: &mut Workspace,
alpha: f32,
a: *const N,
rsa: isize,
csa: isize,
b: *const N,
rsb: isize,
csb: isize,
beta: f32,
c: *mut N,
rsc: isize,
csc: isize,
epi: &E,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
E: Epilogue<MixedGemm<N>>,
{
let epi = *epi;
unsafe {
let (a, rsa, csa, b, rsb, csb) =
prepack_operands::<N>(ws, m, k, n, a, rsa, csa, b, rsb, csb);
debug_assert!(
csa == 1 && rsb == 1,
"small_mn kernel requires A rows / B cols unit-stride along k"
);
let n_row_tiles = m.div_ceil(MT);
let sizeof = core::mem::size_of::<N>();
let bytes = m
.saturating_mul(k)
.saturating_add(k.saturating_mul(n))
.saturating_add(m.saturating_mul(n))
.saturating_mul(sizeof);
let a = Ptr(a as *mut N);
let b = Ptr(b as *mut N);
let c = Ptr(c);
let body = move |q_start: usize, q_end: usize| {
let (a, b, c, epi) = (a, b, c, epi);
let a = a.0 as *const N;
let b = b.0 as *const N;
let c = c.0;
simd.vectorize(|| {
for q in q_start..q_end {
let it = q % n_row_tiles;
let jt = q / n_row_tiles;
let i0 = it * MT;
let j0 = jt * NT;
let mi = core::cmp::min(MT, m - i0);
let nj = core::cmp::min(NT, n - j0);
if mi == MT && nj == NT {
full_tile_mixed::<N, S, E, MT, NT>(
simd, k, i0, j0, alpha, a, rsa, b, csb, beta, c, rsc, csc, &epi,
);
} else {
for cc in 0..nj {
for ir in 0..mi {
cell_dot_mixed::<N, S, E>(
simd,
k,
i0 + ir,
j0 + cc,
alpha,
a,
rsa,
b,
csb,
beta,
c,
rsc,
csc,
&epi,
);
}
}
}
}
});
};
tile_sweep(m, n, bytes, par, body);
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn full_tile_mixed<N, S, E, const MT: usize, const NT: usize>(
simd: S,
k: usize,
i0: usize,
j0: usize,
alpha: f32,
a: *const N,
rsa: isize,
b: *const N,
csb: isize,
beta: f32,
c: *mut N,
rsc: isize,
csc: isize,
epi: &E,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
E: Epilogue<MixedGemm<N>>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let rows: [*const N; MT] = core::array::from_fn(|r| a.offset((i0 + r) as isize * rsa));
let cols: [*const N; NT] = core::array::from_fn(|cc| b.offset((j0 + cc) as isize * csb));
let mut acc: [[<S as SimdOps<f32>>::Reg; MT]; NT] = [[simd.zero(); MT]; NT];
let mut kk = 0;
while kk + lanes <= k {
let av: [<S as SimdOps<f32>>::Reg; MT] =
core::array::from_fn(|r| simd.load_lhs(rows[r].add(kk)));
for cc in 0..NT {
let bv = simd.load_lhs(cols[cc].add(kk));
for r in 0..MT {
acc[cc][r] = simd.mul_add(av[r], bv, acc[cc][r]);
}
}
kk += lanes;
}
for cc in 0..NT {
for r in 0..MT {
let mut dot = simd.reduce_sum(acc[cc][r]);
let mut t = kk;
while t < k {
dot += (*rows[r].add(t)).widen() * (*cols[cc].add(t)).widen();
t += 1;
}
let cp = c.offset((i0 + r) as isize * rsc + (j0 + cc) as isize * csc);
let ov = if beta == 0.0 {
0.0
} else if beta == 1.0 {
(*cp).widen()
} else {
beta * (*cp).widen()
};
let out = alpha * dot + ov;
*cp = if E::IS_IDENTITY {
N::narrow(out)
} else {
epi.apply(out, i0 + r, j0 + cc)
};
}
}
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn cell_dot_mixed<N, S, E>(
simd: S,
k: usize,
i: usize,
j: usize,
alpha: f32,
a: *const N,
rsa: isize,
b: *const N,
csb: isize,
beta: f32,
c: *mut N,
rsc: isize,
csc: isize,
epi: &E,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
E: Epilogue<MixedGemm<N>>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let row = a.offset(i as isize * rsa);
let col = b.offset(j as isize * csb);
let mut acc = simd.zero();
let mut kk = 0;
while kk + lanes <= k {
acc = simd.mul_add(simd.load_lhs(row.add(kk)), simd.load_lhs(col.add(kk)), acc);
kk += lanes;
}
let mut dot = simd.reduce_sum(acc);
while kk < k {
dot += (*row.add(kk)).widen() * (*col.add(kk)).widen();
kk += 1;
}
let cp = c.offset(i as isize * rsc + j as isize * csc);
let ov = if beta == 0.0 {
0.0
} else if beta == 1.0 {
(*cp).widen()
} else {
beta * (*cp).widen()
};
let out = alpha * dot + ov;
*cp = if E::IS_IDENTITY {
N::narrow(out)
} else {
epi.apply(out, i, j)
};
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_int<S>(
simd: S,
m: usize,
k: usize,
n: usize,
par: Parallelism,
ws: &mut Workspace,
alpha: i32,
a: *const i8,
rsa: isize,
csa: isize,
b: *const i8,
rsb: isize,
csb: isize,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
) where
S: KernelSimd<i8, i8, i32, i32>,
{
unsafe {
let (a, rsa, csa, b, rsb, csb) =
prepack_operands::<i8>(ws, m, k, n, a, rsa, csa, b, rsb, csb);
debug_assert!(
csa == 1 && rsb == 1,
"small_mn kernel requires A rows / B cols unit-stride along k"
);
let n_row_tiles = m.div_ceil(MT);
let bytes = m
.saturating_mul(k)
.saturating_add(k.saturating_mul(n))
.saturating_mul(core::mem::size_of::<i8>())
.saturating_add(
m.saturating_mul(n)
.saturating_mul(core::mem::size_of::<i32>()),
);
let a = Ptr(a as *mut i8);
let b = Ptr(b as *mut i8);
let c = Ptr(c);
let body = move |q_start: usize, q_end: usize| {
let (a, b, c) = (a, b, c);
let a = a.0 as *const i8;
let b = b.0 as *const i8;
let c = c.0;
simd.vectorize(|| {
for q in q_start..q_end {
let it = q % n_row_tiles;
let jt = q / n_row_tiles;
let i0 = it * MT;
let j0 = jt * NT;
let mi = core::cmp::min(MT, m - i0);
let nj = core::cmp::min(NT, n - j0);
if mi == MT && nj == NT {
full_tile_int::<S, MT, NT>(
simd, k, i0, j0, alpha, a, rsa, b, csb, beta, c, rsc, csc,
);
} else {
for cc in 0..nj {
for ir in 0..mi {
cell_dot_int::<S>(
simd,
k,
i0 + ir,
j0 + cc,
alpha,
a,
rsa,
b,
csb,
beta,
c,
rsc,
csc,
);
}
}
}
}
});
};
tile_sweep(m, n, bytes, par, body);
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn full_tile_int<S, const MT: usize, const NT: usize>(
simd: S,
k: usize,
i0: usize,
j0: usize,
alpha: i32,
a: *const i8,
rsa: isize,
b: *const i8,
csb: isize,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
) where
S: KernelSimd<i8, i8, i32, i32>,
{
unsafe {
let lanes = <S as SimdOps<i32>>::LANES;
let rows: [*const i8; MT] = core::array::from_fn(|r| a.offset((i0 + r) as isize * rsa));
let cols: [*const i8; NT] = core::array::from_fn(|cc| b.offset((j0 + cc) as isize * csb));
let mut acc: [[<S as SimdOps<i32>>::Reg; MT]; NT] = [[simd.zero(); MT]; NT];
let mut kk = 0;
while kk + lanes <= k {
let av: [<S as SimdOps<i32>>::Reg; MT] = core::array::from_fn(|r| {
<S as KernelSimd<i8, i8, i32, i32>>::load_lhs(simd, rows[r].add(kk))
});
for cc in 0..NT {
let bv = <S as KernelSimd<i8, i8, i32, i32>>::load_lhs(simd, cols[cc].add(kk));
for r in 0..MT {
acc[cc][r] = simd.mul_add(av[r], bv, acc[cc][r]);
}
}
kk += lanes;
}
for cc in 0..NT {
for r in 0..MT {
let mut dot = simd.reduce_sum(acc[cc][r]);
let mut t = kk;
while t < k {
dot = dot.wrapping_add(
(*rows[r].add(t) as i32).wrapping_mul(*cols[cc].add(t) as i32),
);
t += 1;
}
let cp = c.offset((i0 + r) as isize * rsc + (j0 + cc) as isize * csc);
let ov = if beta == 0 {
0
} else if beta == 1 {
*cp
} else {
beta.wrapping_mul(*cp)
};
*cp = alpha.wrapping_mul(dot).wrapping_add(ov);
}
}
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn cell_dot_int<S>(
simd: S,
k: usize,
i: usize,
j: usize,
alpha: i32,
a: *const i8,
rsa: isize,
b: *const i8,
csb: isize,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
) where
S: KernelSimd<i8, i8, i32, i32>,
{
unsafe {
let lanes = <S as SimdOps<i32>>::LANES;
let row = a.offset(i as isize * rsa);
let col = b.offset(j as isize * csb);
let mut acc = simd.zero();
let mut kk = 0;
while kk + lanes <= k {
acc = simd.mul_add(
<S as KernelSimd<i8, i8, i32, i32>>::load_lhs(simd, row.add(kk)),
<S as KernelSimd<i8, i8, i32, i32>>::load_lhs(simd, col.add(kk)),
acc,
);
kk += lanes;
}
let mut dot = simd.reduce_sum(acc);
while kk < k {
dot = dot.wrapping_add((*row.add(kk) as i32).wrapping_mul(*col.add(kk) as i32));
kk += 1;
}
let cp = c.offset(i as isize * rsc + j as isize * csc);
let ov = if beta == 0 {
0
} else if beta == 1 {
*cp
} else {
beta.wrapping_mul(*cp)
};
*cp = alpha.wrapping_mul(dot).wrapping_add(ov);
}
}