use core::marker::PhantomData;
use super::epilogue::Epilogue;
use super::{AlphaStatus, BetaStatus, KernelFamily};
use crate::pack::pack_panels;
use crate::scalar::Float;
use crate::simd::{KernelSimd, SimdOps};
pub struct FloatGemm<T>(PhantomData<T>);
impl<T> Clone for FloatGemm<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for FloatGemm<T> {}
impl<T> KernelFamily for FloatGemm<T>
where
T: Float<Acc = T>,
{
type Lhs = T;
type Rhs = T;
type Acc = T;
type Out = T;
#[inline]
unsafe fn pack_lhs(
dst: *mut T,
src: *const T,
rs: isize,
cs: isize,
mc: usize,
kc: usize,
mr: usize,
) {
unsafe {
pack_panels(
dst, src, rs, cs, mc, kc, mr,
)
}
}
#[inline]
unsafe fn pack_rhs(
dst: *mut T,
src: *const T,
rs: isize,
cs: isize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe {
pack_panels(
dst, src, cs, rs, nc, kc, nr,
)
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn microkernel_epi<S, E, const MR_REG: usize, const NR: usize>(
simd: S,
kc: usize,
alpha: T,
beta: T,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const T,
a_cs: isize,
b: *const T,
b_rs: isize,
b_cs: isize,
c: *mut T,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
last_k: bool,
epi: &E,
scratch: *mut T,
) where
S: KernelSimd<T, T, T, T>,
E: Epilogue<Self>,
{
unsafe {
microkernel_impl::<T, S, E, MR_REG, NR>(
simd,
kc,
alpha,
beta,
alpha_status,
beta_status,
a,
a_cs,
b,
b_rs,
b_cs,
c,
rsc,
csc,
mr_eff,
nr_eff,
row0,
col0,
last_k,
epi,
scratch,
)
}
}
}
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn microkernel_impl<T, S, E, const MR_REG: usize, const NR: usize>(
simd: S,
kc: usize,
alpha: T,
beta: T,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const T,
a_cs: isize,
b: *const T,
b_rs: isize,
b_cs: isize,
c: *mut T,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
last_k: bool,
epi: &E,
scratch: *mut T,
) where
T: Float<Acc = T>,
S: KernelSimd<T, T, T, T>,
E: Epilogue<FloatGemm<T>>,
{
unsafe {
let lanes = <S as SimdOps<T>>::LANES;
let mr = MR_REG * lanes;
let mut acc: [[<S as SimdOps<T>>::Reg; MR_REG]; NR] = [[simd.zero(); MR_REG]; NR];
if nr_eff == NR {
simd.accumulate_tile::<MR_REG, NR>(kc, a, a_cs, b, b_rs, b_cs, &mut acc);
} else {
for p in 0..kc {
let pa = a.offset(p as isize * a_cs);
let a_regs: [<S as SimdOps<T>>::Reg; MR_REG] =
core::array::from_fn(|i| simd.loadu(pa.add(i * lanes)));
let pb = b.offset(p as isize * b_rs);
for j in 0..nr_eff {
let bj = simd.splat(*pb.offset(j as isize * b_cs));
for i in 0..MR_REG {
acc[j][i] = simd.mul_add(a_regs[i], bj, acc[j][i]);
}
}
}
}
if alpha_status == AlphaStatus::Other {
let av = simd.splat(alpha);
for j in 0..NR {
for i in 0..MR_REG {
acc[j][i] = simd.mul(acc[j][i], av);
}
}
}
if (E::IS_IDENTITY || E::VECTOR) && mr_eff == mr && nr_eff == NR && rsc == 1 {
match beta_status {
BetaStatus::Zero => {}
BetaStatus::One => {
for j in 0..NR {
let col = c.offset(j as isize * csc);
for i in 0..MR_REG {
let cv = simd.loadu(col.add(i * lanes));
acc[j][i] = simd.add(cv, acc[j][i]);
}
}
}
BetaStatus::Other => {
let bv = simd.splat(beta);
for j in 0..NR {
let col = c.offset(j as isize * csc);
for i in 0..MR_REG {
let cv = simd.loadu(col.add(i * lanes));
acc[j][i] = simd.mul_add(cv, bv, acc[j][i]);
}
}
}
}
if !E::IS_IDENTITY && last_k {
acc = epi.apply_tile::<S, MR_REG, NR>(simd, acc, row0, col0);
}
for j in 0..NR {
let col = c.offset(j as isize * csc);
for i in 0..MR_REG {
simd.storeu(col.add(i * lanes), acc[j][i]);
}
}
} else {
for j in 0..NR {
for i in 0..MR_REG {
simd.storeu(scratch.add(j * mr + i * lanes), acc[j][i]);
}
}
for j in 0..nr_eff {
for i in 0..mr_eff {
let v = *scratch.add(j * mr + i); let cp = c.offset(i as isize * rsc + j as isize * csc);
let out = match beta_status {
BetaStatus::Zero => v,
BetaStatus::One => *cp + v,
BetaStatus::Other => beta.mul_add(*cp, v), };
*cp = if !E::IS_IDENTITY && last_k {
epi.apply(out, row0 + i, col0 + j)
} else {
out
};
}
}
}
}
}