use core::marker::PhantomData;
use super::epilogue::Epilogue;
use super::{AlphaStatus, BetaStatus, KernelFamily};
use crate::pack::{pack_kgroup_panels, pack_panels};
use crate::scalar::NarrowFloat;
use crate::simd::{KernelSimd, SimdOps};
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn mixed_accumulate<N, S, O, const MR_REG: usize, const NR: usize>(
simd: S,
kc: usize,
a: *const N,
a_cs: isize,
b: *const N,
b_rs: isize,
b_cs: isize,
nr_eff: usize,
acc: &mut [[<S as SimdOps<f32>>::Reg; MR_REG]; NR],
) where
N: NarrowFloat,
O: crate::scalar::Scalar,
S: KernelSimd<N, N, f32, O>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::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<f32>>::Reg; MR_REG] =
core::array::from_fn(|i| simd.load_lhs(pa.add(i * lanes)));
let pb = b.offset(p as isize * b_rs);
for j in 0..NR {
let bj = simd.splat_rhs(*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<f32>>::Reg; MR_REG] =
core::array::from_fn(|i| simd.load_lhs(pa.add(i * lanes)));
let pb = b.offset(p as isize * b_rs);
for j in 0..nr_eff {
let bj = simd.splat_rhs(*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]);
}
}
}
}
}
}
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn mixed_epilogue<Fam, N, S, E, const MR_REG: usize, const NR: usize>(
simd: S,
alpha: f32,
beta: f32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
acc: &mut [[<S as SimdOps<f32>>::Reg; MR_REG]; NR],
c: *mut N,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
epi: &E,
scratch: *mut f32,
) where
N: NarrowFloat,
Fam: KernelFamily<Lhs = N, Rhs = N, Acc = f32, Out = N>,
S: KernelSimd<N, N, f32, N>,
E: Epilogue<Fam>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::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 (E::IS_IDENTITY || E::VECTOR) && 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 {
let r = acc[j][i];
let r = if !E::IS_IDENTITY {
epi.apply_reg(simd, r, row0 + i * lanes, col0 + j)
} else {
r
};
simd.store_out(col.add(i * lanes), r);
}
}
}
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));
let r = simd.add(cv, acc[j][i]);
let r = if !E::IS_IDENTITY {
epi.apply_reg(simd, r, row0 + i * lanes, col0 + j)
} else {
r
};
simd.store_out(col.add(i * lanes), r);
}
}
}
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));
let r = simd.mul_add(cv, bv, acc[j][i]);
let r = if !E::IS_IDENTITY {
epi.apply_reg(simd, r, row0 + i * lanes, col0 + j)
} else {
r
};
simd.store_out(col.add(i * lanes), r);
}
}
}
}
} 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).widen() + v,
BetaStatus::Other => beta * (*cp).widen() + v,
};
*cp = if !E::IS_IDENTITY {
epi.apply(out, row0 + i, col0 + j)
} else {
N::narrow(out)
};
}
}
}
}
}
pub struct MixedGemm<N>(PhantomData<N>);
impl<N> Clone for MixedGemm<N> {
fn clone(&self) -> Self {
*self
}
}
impl<N> Copy for MixedGemm<N> {}
impl<N> KernelFamily for MixedGemm<N>
where
N: NarrowFloat,
{
type Lhs = N;
type Rhs = N;
type Acc = f32;
type Out = N;
const OUT_IS_ACC: bool = false;
#[inline]
unsafe fn pack_lhs(
dst: *mut N,
src: *const N,
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 N,
src: *const N,
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: f32,
beta: f32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const N,
a_cs: isize,
b: *const N,
b_rs: isize,
b_cs: isize,
c: *mut N,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
last_k: bool,
epi: &E,
scratch: *mut f32,
) where
S: KernelSimd<N, N, f32, N>,
E: Epilogue<Self>,
{
debug_assert!(
last_k,
"mixed families are single-panel (kc = k); last_k must be true"
);
let _ = last_k;
unsafe {
let mut acc: [[<S as SimdOps<f32>>::Reg; MR_REG]; NR] = [[simd.zero(); MR_REG]; NR];
mixed_accumulate::<N, S, N, MR_REG, NR>(
simd, kc, a, a_cs, b, b_rs, b_cs, nr_eff, &mut acc,
);
mixed_epilogue::<Self, N, S, E, MR_REG, NR>(
simd,
alpha,
beta,
alpha_status,
beta_status,
&mut acc,
c,
rsc,
csc,
mr_eff,
nr_eff,
row0,
col0,
epi,
scratch,
);
}
}
}
#[derive(Clone, Copy)]
pub struct Bf16DotGemm(PhantomData<()>);
impl Bf16DotGemm {
const Q: usize = 2;
}
impl KernelFamily for Bf16DotGemm {
type Lhs = half::bf16;
type Rhs = half::bf16;
type Acc = f32;
type Out = half::bf16;
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 half::bf16,
src: *const half::bf16,
rs: isize,
cs: isize,
mc: usize,
kc: usize,
mr: usize,
) {
unsafe {
pack_kgroup_panels::<half::bf16, { Self::Q }, _>(dst, src, rs, cs, mc, kc, mr, |v| v)
}
}
#[inline]
unsafe fn pack_rhs(
dst: *mut half::bf16,
src: *const half::bf16,
rs: isize,
cs: isize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe {
pack_kgroup_panels::<half::bf16, { Self::Q }, _>(dst, src, cs, rs, nc, kc, nr, |v| v)
}
}
#[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: f32,
beta: f32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const half::bf16,
_a_cs: isize,
b: *const half::bf16,
_b_rs: isize,
_b_cs: isize,
c: *mut half::bf16,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
last_k: bool,
epi: &E,
scratch: *mut f32,
) where
S: KernelSimd<half::bf16, half::bf16, f32, half::bf16>,
E: Epilogue<Self>,
{
debug_assert!(
last_k,
"mixed families are single-panel (kc = k); last_k must be true"
);
let _ = last_k;
unsafe {
let mut acc: [[<S as SimdOps<f32>>::Reg; MR_REG]; NR] = [[simd.zero(); MR_REG]; NR];
simd.dot_accumulate::<MR_REG, NR>(kc, a, b, &mut acc);
mixed_epilogue::<Self, half::bf16, S, E, MR_REG, NR>(
simd,
alpha,
beta,
alpha_status,
beta_status,
&mut acc,
c,
rsc,
csc,
mr_eff,
nr_eff,
row0,
col0,
epi,
scratch,
);
}
}
}
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn twin_seed<S, const MR_REG: usize, const NR: usize>(
simd: S,
beta_status: BetaStatus,
c: *const f32,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
scratch: *mut f32,
) -> [[<S as SimdOps<f32>>::Reg; MR_REG]; NR]
where
S: SimdOps<f32>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let mr = MR_REG * lanes;
let mut acc: [[<S as SimdOps<f32>>::Reg; MR_REG]; NR] = [[simd.zero(); MR_REG]; NR];
if beta_status == BetaStatus::One {
if mr_eff == mr && nr_eff == NR && rsc == 1 {
for j in 0..NR {
let col = c.offset(j as isize * csc);
for i in 0..MR_REG {
acc[j][i] = simd.loadu(col.add(i * lanes));
}
}
} else {
for x in 0..NR * mr {
*scratch.add(x) = 0.0;
}
for j in 0..nr_eff {
for i in 0..mr_eff {
*scratch.add(j * mr + i) = *c.offset(i as isize * rsc + j as isize * csc);
}
}
for j in 0..NR {
for i in 0..MR_REG {
acc[j][i] = simd.loadu(scratch.add(j * mr + i * lanes));
}
}
}
}
acc
}
}
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
#[inline(always)]
unsafe fn twin_store<S, const MR_REG: usize, const NR: usize>(
simd: S,
acc: &[[<S as SimdOps<f32>>::Reg; MR_REG]; NR],
c: *mut f32,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
scratch: *mut f32,
) where
S: SimdOps<f32>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let mr = MR_REG * lanes;
if mr_eff == mr && nr_eff == NR && rsc == 1 {
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 {
*c.offset(i as isize * rsc + j as isize * csc) = *scratch.add(j * mr + i);
}
}
}
}
}
pub struct MixedGemmF32<N>(PhantomData<N>);
impl<N> Clone for MixedGemmF32<N> {
fn clone(&self) -> Self {
*self
}
}
impl<N> Copy for MixedGemmF32<N> {}
impl<N> KernelFamily for MixedGemmF32<N>
where
N: NarrowFloat,
{
type Lhs = N;
type Rhs = N;
type Acc = f32;
type Out = f32;
#[inline]
unsafe fn pack_lhs(
dst: *mut N,
src: *const N,
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 N,
src: *const N,
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: f32,
beta: f32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const N,
a_cs: isize,
b: *const N,
b_rs: isize,
b_cs: isize,
c: *mut f32,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
last_k: bool,
epi: &E,
scratch: *mut f32,
) where
S: KernelSimd<N, N, f32, f32>,
E: Epilogue<Self>,
{
assert!(E::IS_IDENTITY, "deep-k twin does not fuse epilogues");
debug_assert!(
alpha_status == AlphaStatus::One,
"deep-k twin runs alpha = 1"
);
debug_assert!(
beta_status != BetaStatus::Other,
"deep-k twin runs beta in {{0, 1}}"
);
let _ = (alpha, beta, alpha_status, row0, col0, last_k, epi);
unsafe {
let mut acc =
twin_seed::<S, MR_REG, NR>(simd, beta_status, c, rsc, csc, mr_eff, nr_eff, scratch);
mixed_accumulate::<N, S, f32, MR_REG, NR>(
simd, kc, a, a_cs, b, b_rs, b_cs, nr_eff, &mut acc,
);
twin_store::<S, MR_REG, NR>(simd, &acc, c, rsc, csc, mr_eff, nr_eff, scratch);
}
}
}
#[derive(Clone, Copy)]
pub struct Bf16DotGemmF32(PhantomData<()>);
impl KernelFamily for Bf16DotGemmF32 {
type Lhs = half::bf16;
type Rhs = half::bf16;
type Acc = f32;
type Out = f32;
const FORCE_PACK_LHS: bool = true;
const FORCE_PACK_RHS: bool = true;
const DEPTH_MULTIPLE: usize = 2;
#[inline]
unsafe fn pack_lhs(
dst: *mut half::bf16,
src: *const half::bf16,
rs: isize,
cs: isize,
mc: usize,
kc: usize,
mr: usize,
) {
unsafe { pack_kgroup_panels::<half::bf16, 2, _>(dst, src, rs, cs, mc, kc, mr, |v| v) }
}
#[inline]
unsafe fn pack_rhs(
dst: *mut half::bf16,
src: *const half::bf16,
rs: isize,
cs: isize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe { pack_kgroup_panels::<half::bf16, 2, _>(dst, src, cs, rs, nc, kc, nr, |v| v) }
}
#[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: f32,
beta: f32,
alpha_status: AlphaStatus,
beta_status: BetaStatus,
a: *const half::bf16,
_a_cs: isize,
b: *const half::bf16,
_b_rs: isize,
_b_cs: isize,
c: *mut f32,
rsc: isize,
csc: isize,
mr_eff: usize,
nr_eff: usize,
row0: usize,
col0: usize,
last_k: bool,
epi: &E,
scratch: *mut f32,
) where
S: KernelSimd<half::bf16, half::bf16, f32, f32>,
E: Epilogue<Self>,
{
assert!(E::IS_IDENTITY, "deep-k twin does not fuse epilogues");
debug_assert!(
alpha_status == AlphaStatus::One,
"deep-k twin runs alpha = 1"
);
debug_assert!(
beta_status != BetaStatus::Other,
"deep-k twin runs beta in {{0, 1}}"
);
let _ = (alpha, beta, alpha_status, row0, col0, last_k, epi);
unsafe {
let mut acc =
twin_seed::<S, MR_REG, NR>(simd, beta_status, c, rsc, csc, mr_eff, nr_eff, scratch);
simd.dot_accumulate::<MR_REG, NR>(kc, a, b, &mut acc);
twin_store::<S, MR_REG, NR>(simd, &acc, c, rsc, csc, mr_eff, nr_eff, scratch);
}
}
}