use core::marker::PhantomData;
#[cfg(feature = "epilogue")]
use super::epilogue::{Epilogue, QuantOut};
use super::{AlphaStatus, BetaStatus, KernelFamily};
use crate::pack::{pack_kgroup_panels, pack_panels};
use crate::scalar::Scalar;
use crate::simd::{KernelSimd, SimdOps, VNNI_A_BIAS};
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn int32_epilogue<S, const MR_REG: usize, const NR: usize>(
simd: S,
alpha: i32,
beta: i32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
acc: &mut [[<S as SimdOps<i32>>::Reg; MR_REG]; NR],
c: *mut i32,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
scratch: *mut i32,
) where
S: KernelSimd<i8, i8, i32, i32>,
{
unsafe {
let lanes = <S as SimdOps<i32>>::LANES;
let mr = MR_REG * lanes;
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 mr_eff == mr && nr_eff == NR && rsc == 1 {
match beta_status {
BetaStatus::Zero => {
for j in 0..NR {
let col = c.offset(j as isize * csc);
for i in 0..MR_REG {
simd.store_out(col.add(i * lanes), acc[j][i]);
}
}
}
BetaStatus::One => {
for j in 0..NR {
let col = c.offset(j as isize * csc);
for i in 0..MR_REG {
let cv = simd.load_out(col.add(i * lanes));
simd.store_out(col.add(i * lanes), 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.load_out(col.add(i * lanes));
simd.store_out(col.add(i * lanes), simd.mul_add(cv, bv, 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).wrapping_add(v),
BetaStatus::Other => beta.wrapping_mul(*cp).wrapping_add(v),
};
*cp = out;
}
}
}
}
}
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn i32_accumulate<S, O, const MR_REG: usize, const NR: usize>(
simd: S,
kc: usize,
a: *const i8,
a_cs: isize,
b: *const i8,
b_rs: isize,
b_cs: isize,
nr_eff: usize,
acc: &mut [[<S as SimdOps<i32>>::Reg; MR_REG]; NR],
) where
O: Scalar,
S: KernelSimd<i8, i8, i32, O>,
{
unsafe {
let lanes = <S as SimdOps<i32>>::LANES;
if nr_eff == NR {
for p in 0..kc {
let pa = a.offset(p as isize * a_cs);
let a_regs: [<S as SimdOps<i32>>::Reg; MR_REG] = core::array::from_fn(|i| {
<S as KernelSimd<i8, i8, i32, O>>::load_lhs(simd, pa.add(i * lanes))
});
let pb = b.offset(p as isize * b_rs);
for j in 0..NR {
let bj = <S as KernelSimd<i8, i8, i32, O>>::splat_rhs(
simd,
*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]);
}
}
}
} else {
for p in 0..kc {
let pa = a.offset(p as isize * a_cs);
let a_regs: [<S as SimdOps<i32>>::Reg; MR_REG] = core::array::from_fn(|i| {
<S as KernelSimd<i8, i8, i32, O>>::load_lhs(simd, pa.add(i * lanes))
});
let pb = b.offset(p as isize * b_rs);
for j in 0..nr_eff {
let bj = <S as KernelSimd<i8, i8, i32, O>>::splat_rhs(
simd,
*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]);
}
}
}
}
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn requant_scratch_epilogue<F, S, E, O, const MR_REG: usize, const NR: usize>(
simd: S,
acc: &[[<S as SimdOps<i32>>::Reg; MR_REG]; NR],
c: *mut O,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
epi: &E,
scratch: *mut i32,
) where
O: QuantOut,
F: KernelFamily<Acc = i32, Out = O>,
S: KernelSimd<F::Lhs, F::Rhs, i32, O>,
E: Epilogue<F>,
{
unsafe {
let lanes = <S as SimdOps<i32>>::LANES;
let mr = MR_REG * lanes;
for j in 0..NR {
for i in 0..MR_REG {
simd.storeu(scratch.add(j * mr + i * lanes), acc[j][i]);
}
}
if E::VECTOR_STORE && <S as KernelSimd<F::Lhs, F::Rhs, i32, O>>::REQUANT_VECTOR && rsc == 1
{
for j in 0..nr_eff {
let src_col = scratch.add(j * mr);
let dst_col = c.offset(j as isize * csc); let mut i = 0;
while i + lanes <= mr_eff {
epi.apply_store(simd, src_col.add(i), dst_col.add(i), row0 + i, col0 + j);
i += lanes;
}
for i in i..mr_eff {
*dst_col.add(i) = epi.apply(*src_col.add(i), row0 + i, col0 + j);
}
}
} else {
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);
*cp = epi.apply(v, row0 + i, col0 + j);
}
}
}
}
}
#[derive(Clone, Copy)]
pub struct IntGemm(PhantomData<()>);
impl KernelFamily for IntGemm {
type Lhs = i8;
type Rhs = i8;
type Acc = i32;
type Out = i32;
#[inline]
unsafe fn pack_lhs(
dst: *mut i8,
src: *const i8,
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 i8,
src: *const i8,
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, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn microkernel<S, const MR_REG: usize, const NR: usize>(
simd: S,
kc: usize,
alpha: i32,
beta: i32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const i8,
a_cs: isize,
b: *const i8,
b_rs: isize,
b_cs: isize,
c: *mut i32,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
scratch: *mut i32,
) where
S: KernelSimd<i8, i8, i32, i32>,
{
unsafe {
let mut acc: [[<S as SimdOps<i32>>::Reg; MR_REG]; NR] = [[simd.zero(); MR_REG]; NR];
i32_accumulate::<S, i32, MR_REG, NR>(
simd, kc, a, a_cs, b, b_rs, b_cs, nr_eff, &mut acc,
);
int32_epilogue::<S, MR_REG, NR>(
simd,
alpha,
beta,
alpha_status,
beta_status,
&mut acc,
c,
rsc,
csc,
mr_eff,
nr_eff,
scratch,
);
}
}
}
#[inline(always)]
pub(crate) fn vnni_a_xform(v: i8) -> i8 {
((v as i32 + VNNI_A_BIAS) as u8) as i8
}
#[derive(Clone, Copy)]
pub struct IntGemmVnni(PhantomData<()>);
impl IntGemmVnni {
const Q: usize = 4;
}
impl KernelFamily for IntGemmVnni {
type Lhs = i8;
type Rhs = i8;
type Acc = i32;
type Out = i32;
const FORCE_PACK_LHS: bool = true;
const FORCE_PACK_RHS: bool = true;
const DEPTH_MULTIPLE: usize = Self::Q;
#[inline]
unsafe fn pack_lhs(
dst: *mut i8,
src: *const i8,
rs: isize,
cs: isize,
mc: usize,
kc: usize,
mr: usize,
) {
unsafe {
pack_kgroup_panels::<i8, { Self::Q }, _>(dst, src, rs, cs, mc, kc, mr, vnni_a_xform)
}
}
#[inline]
unsafe fn pack_rhs(
dst: *mut i8,
src: *const i8,
rs: isize,
cs: isize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe { pack_kgroup_panels::<i8, { Self::Q }, _>(dst, src, cs, rs, nc, kc, nr, |v| v) }
}
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn microkernel<S, const MR_REG: usize, const NR: usize>(
simd: S,
kc: usize,
alpha: i32,
beta: i32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const i8,
_a_cs: isize,
b: *const i8,
_b_rs: isize,
_b_cs: isize,
c: *mut i32,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
scratch: *mut i32,
) where
S: KernelSimd<i8, i8, i32, i32>,
{
unsafe {
let mut acc: [[<S as SimdOps<i32>>::Reg; MR_REG]; NR] = [[simd.zero(); MR_REG]; NR];
<S as KernelSimd<i8, i8, i32, i32>>::dot_accumulate::<MR_REG, NR>(
simd, kc, a, b, &mut acc,
);
int32_epilogue::<S, MR_REG, NR>(
simd,
alpha,
beta,
alpha_status,
beta_status,
&mut acc,
c,
rsc,
csc,
mr_eff,
nr_eff,
scratch,
);
}
}
}
#[cfg(feature = "epilogue")]
#[derive(Clone, Copy)]
pub struct IntGemmQ<O = i8>(PhantomData<O>);
#[cfg(feature = "epilogue")]
impl<O: QuantOut> KernelFamily for IntGemmQ<O> {
type Lhs = i8;
type Rhs = i8;
type Acc = i32;
type Out = O;
const OUT_IS_ACC: bool = false;
#[inline]
unsafe fn pack_lhs(
dst: *mut i8,
src: *const i8,
rs: isize,
cs: isize,
mc: usize,
kc: usize,
mr: usize,
) {
unsafe { <IntGemm as KernelFamily>::pack_lhs(dst, src, rs, cs, mc, kc, mr) }
}
#[inline]
unsafe fn pack_rhs(
dst: *mut i8,
src: *const i8,
rs: isize,
cs: isize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe { <IntGemm as KernelFamily>::pack_rhs(dst, src, rs, cs, kc, nc, 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: i32,
_beta: i32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const i8,
a_cs: isize,
b: *const i8,
b_rs: isize,
b_cs: isize,
c: *mut O,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
_last_k: bool,
epi: &E,
scratch: *mut i32,
) where
S: KernelSimd<i8, i8, i32, O>,
E: Epilogue<Self>,
{
debug_assert!(matches!(beta_status, BetaStatus::Zero));
debug_assert!(matches!(alpha_status, AlphaStatus::One));
unsafe {
let mut acc: [[<S as SimdOps<i32>>::Reg; MR_REG]; NR] = [[simd.zero(); MR_REG]; NR];
i32_accumulate::<S, O, MR_REG, NR>(simd, kc, a, a_cs, b, b_rs, b_cs, nr_eff, &mut acc);
requant_scratch_epilogue::<Self, S, E, O, MR_REG, NR>(
simd, &acc, c, rsc, csc, mr_eff, nr_eff, row0, col0, epi, scratch,
);
}
}
}
#[cfg(feature = "epilogue")]
#[derive(Clone, Copy)]
pub struct IntGemmVnniQ<O = i8>(PhantomData<O>);
#[cfg(feature = "epilogue")]
impl<O> IntGemmVnniQ<O> {
const Q: usize = 4;
}
#[cfg(feature = "epilogue")]
impl<O: QuantOut> KernelFamily for IntGemmVnniQ<O> {
type Lhs = i8;
type Rhs = i8;
type Acc = i32;
type Out = O;
const OUT_IS_ACC: bool = false;
const FORCE_PACK_LHS: bool = true;
const FORCE_PACK_RHS: bool = true;
const DEPTH_MULTIPLE: usize = Self::Q;
#[inline]
unsafe fn pack_lhs(
dst: *mut i8,
src: *const i8,
rs: isize,
cs: isize,
mc: usize,
kc: usize,
mr: usize,
) {
unsafe { <IntGemmVnni as KernelFamily>::pack_lhs(dst, src, rs, cs, mc, kc, mr) }
}
#[inline]
unsafe fn pack_rhs(
dst: *mut i8,
src: *const i8,
rs: isize,
cs: isize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe { <IntGemmVnni as KernelFamily>::pack_rhs(dst, src, rs, cs, kc, nc, 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: i32,
_beta: i32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const i8,
_a_cs: isize,
b: *const i8,
_b_rs: isize,
_b_cs: isize,
c: *mut O,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
_last_k: bool,
epi: &E,
scratch: *mut i32,
) where
S: KernelSimd<i8, i8, i32, O>,
E: Epilogue<Self>,
{
debug_assert!(matches!(beta_status, BetaStatus::Zero));
debug_assert!(matches!(alpha_status, AlphaStatus::One));
unsafe {
let mut acc: [[<S as SimdOps<i32>>::Reg; MR_REG]; NR] = [[simd.zero(); MR_REG]; NR];
<S as KernelSimd<i8, i8, i32, O>>::dot_accumulate::<MR_REG, NR>(
simd, kc, a, b, &mut acc,
);
requant_scratch_epilogue::<Self, S, E, O, MR_REG, NR>(
simd, &acc, c, rsc, csc, mr_eff, nr_eff, row0, col0, epi, scratch,
);
}
}
}