#[cfg(feature = "std")]
use std::sync::OnceLock;
use super::float::Dispatched;
#[cfg(feature = "epilogue")]
use super::float::FusedScalar;
use super::isa::{ForcedIsa, forced_isa};
use super::{
GemmScalar, PackedConsume, Task, orient_transpose, small_mn_eligible, small_mn_pack_eligible,
};
use crate::driver::{self, alpha_status, beta_status};
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
use crate::kernel::Bf16DotGemm;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
use crate::kernel::Bf16DotGemmF32;
use crate::kernel::KernelFamily;
use crate::kernel::MixedGemmF32;
#[cfg(feature = "epilogue")]
use crate::kernel::epilogue::FusedEpi;
use crate::kernel::epilogue::{Epilogue, Identity};
use crate::kernel::{AlphaStatus, BetaStatus, MixedGemm};
use crate::parallel::Parallelism;
use crate::scalar::NarrowFloat;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
use crate::simd::Avx512Bf16;
#[cfg(target_arch = "aarch64")]
use crate::simd::Neon;
use crate::simd::ScalarTok;
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
use crate::simd::Simd128;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
use crate::simd::{Avx512F, Fma};
use crate::simd::{KernelSimd, SimdOps};
use crate::special::{gemv, small_k, small_mn};
use crate::tuning;
use crate::workspace::Workspace;
use half::{bf16, f16};
#[cfg(feature = "half")]
unsafe fn scale_c_narrow<N: NarrowFloat>(
beta: N,
c: *mut N,
m: usize,
n: usize,
rsc: isize,
csc: isize,
) {
unsafe {
let b = beta.widen();
for j in 0..n {
for i in 0..m {
let p = c.offset(i as isize * rsc + j as isize * csc);
if beta == N::ZERO {
*p = N::ZERO;
} else if beta != N::ONE {
*p = N::narrow(b * (*p).widen());
}
}
}
}
}
#[cfg(feature = "half")]
trait DeepKTwin: KernelFamily {
type Twin: KernelFamily<Lhs = Self::Lhs, Rhs = Self::Rhs, Acc = Self::Acc, Out = f32>;
}
#[cfg(feature = "half")]
impl<N: NarrowFloat> DeepKTwin for MixedGemm<N> {
type Twin = MixedGemmF32<N>;
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
impl DeepKTwin for Bf16DotGemm {
type Twin = Bf16DotGemmF32;
}
#[cfg(feature = "half")]
#[inline]
unsafe fn run_deep_k_twin<N, Tw, S, const MR_REG: usize, const NR: usize>(
simd: S,
t: &Task<N>,
par: Parallelism,
ws: &mut Workspace,
) where
N: NarrowFloat,
Tw: KernelFamily<Lhs = N, Rhs = N, Acc = f32, Out = f32>,
S: KernelSimd<N, N, f32, N> + KernelSimd<N, N, f32, f32>,
{
unsafe {
let (m, n, k) = (t.m, t.n, t.k);
let mut scratch_ws = Workspace::new();
let scratch = scratch_ws.regions::<f32>(m.saturating_mul(n), 1, 0).a_base;
driver::run::<Tw, S, MR_REG, NR>(
simd, m, k, n, 1.0f32, t.a, t.rsa, t.csa, t.b, t.rsb, t.csb, 0.0f32, scratch, 1,
m as isize, par, ws,
);
let alpha = t.alpha.widen();
let beta = t.beta.widen();
let ash = alpha_status(alpha);
let bst = beta_status(beta);
let (c, rsc, csc) = (t.c, t.rsc, t.csc);
simd.vectorize(|| {
let lanes = <S as SimdOps<f32>>::LANES;
let av = simd.splat(alpha);
let bv = simd.splat(beta);
for j in 0..n {
let sc = scratch.add(j * m); let cc = c.offset(j as isize * csc); if rsc == 1 {
let mut i = 0;
while i + lanes <= m {
let mut r = simd.loadu(sc.add(i));
if ash == AlphaStatus::Other {
r = simd.mul(r, av);
}
r = match bst {
BetaStatus::Zero => r,
BetaStatus::One => {
let cv = <S as KernelSimd<N, N, f32, N>>::load_out(simd, cc.add(i));
simd.add(cv, r)
}
BetaStatus::Other => {
let cv = <S as KernelSimd<N, N, f32, N>>::load_out(simd, cc.add(i));
simd.mul_add(cv, bv, r)
}
};
<S as KernelSimd<N, N, f32, N>>::store_out(simd, cc.add(i), r);
i += lanes;
}
while i < m {
let mut r = *sc.add(i);
if ash == AlphaStatus::Other {
r *= alpha;
}
r = match bst {
BetaStatus::Zero => r,
BetaStatus::One => (*cc.add(i)).widen() + r,
BetaStatus::Other => beta * (*cc.add(i)).widen() + r,
};
*cc.add(i) = N::narrow(r);
i += 1;
}
} else {
for i in 0..m {
let cp = cc.offset(i as isize * rsc);
let mut r = *sc.add(i);
if ash == AlphaStatus::Other {
r *= alpha;
}
r = match bst {
BetaStatus::Zero => r,
BetaStatus::One => (*cp).widen() + r,
BetaStatus::Other => beta * (*cp).widen() + r,
};
*cp = N::narrow(r);
}
}
}
});
}
}
#[cfg(feature = "half")]
#[inline]
unsafe fn run_typed_mixed_epi<N, Fam, S, const MR_REG: usize, const NR: usize, E>(
simd: S,
mut t: Task<N>,
mut epi: E,
par: Parallelism,
ws: &mut Workspace,
) where
N: NarrowFloat,
Fam: KernelFamily<Lhs = N, Rhs = N, Acc = f32, Out = N> + DeepKTwin,
S: KernelSimd<N, N, f32, N> + KernelSimd<N, N, f32, f32>,
E: Epilogue<Fam> + Epilogue<MixedGemm<N>>,
{
unsafe {
if <E as Epilogue<Fam>>::IS_IDENTITY
&& (t.n == 1 || t.m == 1)
&& core::cmp::min(t.m, t.n) <= tuning::gemv_threshold()
{
gemv::run_mixed::<N, S>(
simd,
t.m,
t.k,
t.n,
par,
t.alpha.widen(),
t.a,
t.rsa,
t.csa,
t.b,
t.rsb,
t.csb,
t.beta.widen(),
t.c,
t.rsc,
t.csc,
);
return;
}
if orient_transpose(&mut t) {
<E as Epilogue<Fam>>::on_orient_swap(&mut epi);
}
if small_mn_eligible(&t) || small_mn_pack_eligible(&t) {
small_mn::run_mixed_epi::<N, S, E>(
simd,
t.m,
t.k,
t.n,
par,
ws,
t.alpha.widen(),
t.a,
t.rsa,
t.csa,
t.b,
t.rsb,
t.csb,
t.beta.widen(),
t.c,
t.rsc,
t.csc,
&epi,
);
return;
}
if t.k <= tuning::small_k_threshold() {
small_k::run_epi::<MixedGemm<N>, S, E, MR_REG, NR>(
simd,
t.m,
t.k,
t.n,
t.alpha.widen(),
t.a,
t.rsa,
t.csa,
t.b,
t.rsb,
t.csb,
t.beta.widen(),
t.c,
t.rsc,
t.csc,
&epi,
par,
ws,
);
return;
}
if <E as Epilogue<Fam>>::IS_IDENTITY {
let engage_deep_k = NR
.checked_mul(t.k)
.and_then(|x| x.checked_mul(core::mem::size_of::<N>()))
.is_some_and(|bytes| bytes > crate::cache::deep_k_engage_bytes());
if engage_deep_k {
run_deep_k_twin::<N, Fam::Twin, S, MR_REG, NR>(simd, &t, par, ws);
return;
}
}
driver::run_epilogue::<Fam, S, E, MR_REG, NR>(
simd,
t.m,
t.k,
t.n,
t.alpha.widen(),
t.a,
t.rsa,
t.csa,
t.b,
t.rsb,
t.csb,
t.beta.widen(),
t.c,
t.rsc,
t.csc,
&epi,
par,
ws,
);
}
}
#[cfg(feature = "half")]
#[inline]
unsafe fn run_typed_mixed<N, Fam, S, const MR_REG: usize, const NR: usize>(
simd: S,
t: Task<N>,
par: Parallelism,
ws: &mut Workspace,
) where
N: NarrowFloat,
Fam: KernelFamily<Lhs = N, Rhs = N, Acc = f32, Out = N> + DeepKTwin,
S: KernelSimd<N, N, f32, N> + KernelSimd<N, N, f32, f32>,
{
unsafe { run_typed_mixed_epi::<N, Fam, S, MR_REG, NR, Identity>(simd, t, Identity, par, ws) }
}
#[cfg(all(feature = "half", feature = "epilogue"))]
#[inline]
unsafe fn run_typed_mixed_fused<N, Fam, S, const MR_REG: usize, const NR: usize>(
simd: S,
t: Task<N>,
epi: FusedEpi<N>,
par: Parallelism,
ws: &mut Workspace,
) where
N: NarrowFloat,
Fam: KernelFamily<Lhs = N, Rhs = N, Acc = f32, Out = N> + DeepKTwin,
S: KernelSimd<N, N, f32, N> + KernelSimd<N, N, f32, f32>,
FusedEpi<N>: Epilogue<Fam> + Epilogue<MixedGemm<N>>,
{
unsafe { run_typed_mixed_epi::<N, Fam, S, MR_REG, NR, FusedEpi<N>>(simd, t, epi, par, ws) }
}
#[cfg(all(feature = "half", feature = "epilogue"))]
unsafe fn fused_degenerate_mixed<N>(t: &Task<N>, epi: &FusedEpi<N>)
where
N: NarrowFloat,
FusedEpi<N>: Epilogue<MixedGemm<N>>,
{
unsafe {
for j in 0..t.n {
for i in 0..t.m {
let p = t.c.offset(i as isize * t.rsc + j as isize * t.csc);
let base: f32 = if t.beta == N::ZERO {
0.0
} else if t.beta == N::ONE {
(*p).widen()
} else {
t.beta.widen() * (*p).widen()
};
*p = <FusedEpi<N> as Epilogue<MixedGemm<N>>>::apply(epi, base, i, j);
}
}
}
}
#[cfg(feature = "half")]
#[inline]
unsafe fn run_packed_typed_mixed<N, Fam, S, const MR_REG: usize, const NR: usize>(
simd: S,
req: PackedConsume<N>,
par: Parallelism,
ws: &mut Workspace,
) where
N: NarrowFloat,
Fam: KernelFamily<Lhs = N, Rhs = N, Acc = f32, Out = N>,
S: KernelSimd<N, N, f32, N>,
{
unsafe {
debug_assert_eq!(NR, req.nr, "prepacked RHS panel width != kernel NR");
driver::run_packed_rhs::<Fam, S, MR_REG, NR>(
simd,
req.m,
req.k,
req.n,
req.alpha.widen(),
req.a,
req.rsa,
req.csa,
req.packed,
req.kc,
req.nc,
req.beta.widen(),
req.c,
req.rsc,
req.csc,
par,
ws,
);
}
}
#[cfg(all(feature = "half", feature = "epilogue"))]
#[inline]
unsafe fn run_typed_mixed_packed_fused<N, Fam, S, const MR_REG: usize, const NR: usize>(
simd: S,
req: PackedConsume<N>,
epi: FusedEpi<N>,
par: Parallelism,
ws: &mut Workspace,
) where
N: NarrowFloat,
Fam: KernelFamily<Lhs = N, Rhs = N, Acc = f32, Out = N>,
S: KernelSimd<N, N, f32, N>,
FusedEpi<N>: Epilogue<Fam>,
{
unsafe {
debug_assert_eq!(NR, req.nr, "prepacked RHS panel width != kernel NR");
driver::run_packed_rhs_epilogue::<Fam, S, FusedEpi<N>, MR_REG, NR>(
simd,
req.m,
req.k,
req.n,
req.alpha.widen(),
req.a,
req.rsa,
req.csa,
req.packed,
req.kc,
req.nc,
req.beta.widen(),
req.c,
req.rsc,
req.csc,
&epi,
par,
ws,
);
}
}
#[cfg(feature = "half")]
unsafe fn gemm_f16_scalar(t: Task<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<f16, MixedGemm<f16>, ScalarTok, 4, 4>(ScalarTok, t, par, ws) }
}
#[cfg(feature = "half")]
unsafe fn gemm_bf16_scalar(t: Task<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<bf16, MixedGemm<bf16>, ScalarTok, 4, 4>(ScalarTok, t, par, ws) }
}
#[cfg(feature = "half")]
unsafe fn gemm_f16_scalar_packed(r: PackedConsume<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed_mixed::<f16, MixedGemm<f16>, ScalarTok, 4, 4>(ScalarTok, r, par, ws) }
}
#[cfg(feature = "half")]
unsafe fn gemm_bf16_scalar_packed(r: PackedConsume<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe {
run_packed_typed_mixed::<bf16, MixedGemm<bf16>, ScalarTok, 4, 4>(ScalarTok, r, par, ws)
}
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f16_fma(t: Task<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<f16, MixedGemm<f16>, Fma, 2, 6>(Fma, t, par, ws) }
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_bf16_fma(t: Task<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<bf16, MixedGemm<bf16>, Fma, 2, 6>(Fma, t, par, ws) }
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f16_fma_packed(r: PackedConsume<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed_mixed::<f16, MixedGemm<f16>, Fma, 2, 6>(Fma, r, par, ws) }
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_bf16_fma_packed(r: PackedConsume<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed_mixed::<bf16, MixedGemm<bf16>, Fma, 2, 6>(Fma, r, par, ws) }
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f16_avx512f(t: Task<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<f16, MixedGemm<f16>, Avx512F, 2, 12>(Avx512F, t, par, ws) }
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_bf16_avx512f(t: Task<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<bf16, MixedGemm<bf16>, Avx512F, 2, 12>(Avx512F, t, par, ws) }
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f16_avx512f_packed(r: PackedConsume<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed_mixed::<f16, MixedGemm<f16>, Avx512F, 2, 12>(Avx512F, r, par, ws) }
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_bf16_avx512f_packed(r: PackedConsume<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed_mixed::<bf16, MixedGemm<bf16>, Avx512F, 2, 12>(Avx512F, r, par, ws) }
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_bf16_avx512bf16(t: Task<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<bf16, Bf16DotGemm, Avx512Bf16, 2, 12>(Avx512Bf16, t, par, ws) }
}
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_bf16_avx512bf16_packed(
r: PackedConsume<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_packed_typed_mixed::<bf16, Bf16DotGemm, Avx512Bf16, 2, 12>(Avx512Bf16, r, par, ws)
}
}
#[cfg(all(feature = "half", target_arch = "aarch64"))]
unsafe fn gemm_f16_neon(t: Task<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<f16, MixedGemm<f16>, Neon, 4, 4>(Neon, t, par, ws) }
}
#[cfg(all(feature = "half", target_arch = "aarch64"))]
unsafe fn gemm_bf16_neon(t: Task<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<bf16, MixedGemm<bf16>, Neon, 4, 4>(Neon, t, par, ws) }
}
#[cfg(all(feature = "half", target_arch = "aarch64"))]
unsafe fn gemm_f16_neon_packed(r: PackedConsume<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed_mixed::<f16, MixedGemm<f16>, Neon, 4, 4>(Neon, r, par, ws) }
}
#[cfg(all(feature = "half", target_arch = "aarch64"))]
unsafe fn gemm_bf16_neon_packed(r: PackedConsume<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed_mixed::<bf16, MixedGemm<bf16>, Neon, 4, 4>(Neon, r, par, ws) }
}
#[cfg(all(feature = "half", target_arch = "wasm32", target_feature = "simd128"))]
unsafe fn gemm_f16_simd128(t: Task<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<f16, MixedGemm<f16>, Simd128, 2, 4>(Simd128, t, par, ws) }
}
#[cfg(all(feature = "half", target_arch = "wasm32", target_feature = "simd128"))]
unsafe fn gemm_bf16_simd128(t: Task<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed_mixed::<bf16, MixedGemm<bf16>, Simd128, 2, 4>(Simd128, t, par, ws) }
}
#[cfg(all(feature = "half", target_arch = "wasm32", target_feature = "simd128"))]
unsafe fn gemm_f16_simd128_packed(r: PackedConsume<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed_mixed::<f16, MixedGemm<f16>, Simd128, 2, 4>(Simd128, r, par, ws) }
}
#[cfg(all(feature = "half", target_arch = "wasm32", target_feature = "simd128"))]
unsafe fn gemm_bf16_simd128_packed(r: PackedConsume<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed_mixed::<bf16, MixedGemm<bf16>, Simd128, 2, 4>(Simd128, r, par, ws) }
}
#[cfg(all(feature = "half", feature = "epilogue"))]
unsafe fn gemm_f16_scalar_fused(
t: Task<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_fused::<f16, MixedGemm<f16>, ScalarTok, 4, 4>(ScalarTok, t, epi, par, ws)
}
}
#[cfg(all(feature = "half", feature = "epilogue"))]
unsafe fn gemm_bf16_scalar_fused(
t: Task<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_fused::<bf16, MixedGemm<bf16>, ScalarTok, 4, 4>(ScalarTok, t, epi, par, ws)
}
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_f16_fma_fused(
t: Task<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_mixed_fused::<f16, MixedGemm<f16>, Fma, 2, 6>(Fma, t, epi, par, ws) }
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_bf16_fma_fused(
t: Task<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_mixed_fused::<bf16, MixedGemm<bf16>, Fma, 2, 6>(Fma, t, epi, par, ws) }
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_f16_avx512f_fused(
t: Task<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_fused::<f16, MixedGemm<f16>, Avx512F, 2, 12>(Avx512F, t, epi, par, ws)
}
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_bf16_avx512f_fused(
t: Task<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_fused::<bf16, MixedGemm<bf16>, Avx512F, 2, 12>(Avx512F, t, epi, par, ws)
}
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_bf16_avx512bf16_fused(
t: Task<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_fused::<bf16, Bf16DotGemm, Avx512Bf16, 2, 12>(Avx512Bf16, t, epi, par, ws)
}
}
#[cfg(all(feature = "half", feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_f16_neon_fused(
t: Task<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_mixed_fused::<f16, MixedGemm<f16>, Neon, 4, 4>(Neon, t, epi, par, ws) }
}
#[cfg(all(feature = "half", feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_bf16_neon_fused(
t: Task<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_mixed_fused::<bf16, MixedGemm<bf16>, Neon, 4, 4>(Neon, t, epi, par, ws) }
}
#[cfg(all(
feature = "half",
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_f16_simd128_fused(
t: Task<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_mixed_fused::<f16, MixedGemm<f16>, Simd128, 2, 4>(Simd128, t, epi, par, ws) }
}
#[cfg(all(
feature = "half",
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_bf16_simd128_fused(
t: Task<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_fused::<bf16, MixedGemm<bf16>, Simd128, 2, 4>(Simd128, t, epi, par, ws)
}
}
#[cfg(all(feature = "half", feature = "epilogue"))]
unsafe fn gemm_f16_scalar_packed_fused(
r: PackedConsume<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<f16, MixedGemm<f16>, ScalarTok, 4, 4>(
ScalarTok, r, epi, par, ws,
)
}
}
#[cfg(all(feature = "half", feature = "epilogue"))]
unsafe fn gemm_bf16_scalar_packed_fused(
r: PackedConsume<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<bf16, MixedGemm<bf16>, ScalarTok, 4, 4>(
ScalarTok, r, epi, par, ws,
)
}
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_f16_fma_packed_fused(
r: PackedConsume<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_mixed_packed_fused::<f16, MixedGemm<f16>, Fma, 2, 6>(Fma, r, epi, par, ws) }
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_bf16_fma_packed_fused(
r: PackedConsume<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<bf16, MixedGemm<bf16>, Fma, 2, 6>(Fma, r, epi, par, ws)
}
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_f16_avx512f_packed_fused(
r: PackedConsume<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<f16, MixedGemm<f16>, Avx512F, 2, 12>(
Avx512F, r, epi, par, ws,
)
}
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_bf16_avx512f_packed_fused(
r: PackedConsume<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<bf16, MixedGemm<bf16>, Avx512F, 2, 12>(
Avx512F, r, epi, par, ws,
)
}
}
#[cfg(all(
feature = "half",
feature = "epilogue",
any(target_arch = "x86", target_arch = "x86_64")
))]
unsafe fn gemm_bf16_avx512bf16_packed_fused(
r: PackedConsume<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<bf16, Bf16DotGemm, Avx512Bf16, 2, 12>(
Avx512Bf16, r, epi, par, ws,
)
}
}
#[cfg(all(feature = "half", feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_f16_neon_packed_fused(
r: PackedConsume<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<f16, MixedGemm<f16>, Neon, 4, 4>(Neon, r, epi, par, ws)
}
}
#[cfg(all(feature = "half", feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_bf16_neon_packed_fused(
r: PackedConsume<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<bf16, MixedGemm<bf16>, Neon, 4, 4>(Neon, r, epi, par, ws)
}
}
#[cfg(all(
feature = "half",
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_f16_simd128_packed_fused(
r: PackedConsume<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<f16, MixedGemm<f16>, Simd128, 2, 4>(Simd128, r, epi, par, ws)
}
}
#[cfg(all(
feature = "half",
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_bf16_simd128_packed_fused(
r: PackedConsume<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
run_typed_mixed_packed_fused::<bf16, MixedGemm<bf16>, Simd128, 2, 4>(
Simd128, r, epi, par, ws,
)
}
}
#[cfg(feature = "half")]
const DISP_F16_SCALAR: Dispatched<f16> = Dispatched {
run: gemm_f16_scalar,
run_packed: gemm_f16_scalar_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f16_scalar_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f16_scalar_packed_fused,
mr: 4,
nr: 4,
depth_multiple: 1,
};
#[cfg(feature = "half")]
const DISP_BF16_SCALAR: Dispatched<bf16> = Dispatched {
run: gemm_bf16_scalar,
run_packed: gemm_bf16_scalar_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_bf16_scalar_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_bf16_scalar_packed_fused,
mr: 4,
nr: 4,
depth_multiple: 1,
};
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
const DISP_F16_FMA: Dispatched<f16> = Dispatched {
run: gemm_f16_fma,
run_packed: gemm_f16_fma_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f16_fma_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f16_fma_packed_fused,
mr: 16,
nr: 6,
depth_multiple: 1,
};
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
const DISP_BF16_FMA: Dispatched<bf16> = Dispatched {
run: gemm_bf16_fma,
run_packed: gemm_bf16_fma_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_bf16_fma_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_bf16_fma_packed_fused,
mr: 16,
nr: 6,
depth_multiple: 1,
};
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
const DISP_F16_AVX512F: Dispatched<f16> = Dispatched {
run: gemm_f16_avx512f,
run_packed: gemm_f16_avx512f_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f16_avx512f_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f16_avx512f_packed_fused,
mr: 32,
nr: 12,
depth_multiple: 1,
};
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
const DISP_BF16_AVX512F: Dispatched<bf16> = Dispatched {
run: gemm_bf16_avx512f,
run_packed: gemm_bf16_avx512f_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_bf16_avx512f_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_bf16_avx512f_packed_fused,
mr: 32,
nr: 12,
depth_multiple: 1,
};
#[cfg(all(feature = "half", any(target_arch = "x86", target_arch = "x86_64")))]
const DISP_BF16_AVX512BF16: Dispatched<bf16> = Dispatched {
run: gemm_bf16_avx512bf16,
run_packed: gemm_bf16_avx512bf16_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_bf16_avx512bf16_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_bf16_avx512bf16_packed_fused,
mr: 32,
nr: 12,
depth_multiple: 2,
};
#[cfg(all(feature = "half", target_arch = "aarch64"))]
const DISP_F16_NEON: Dispatched<f16> = Dispatched {
run: gemm_f16_neon,
run_packed: gemm_f16_neon_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f16_neon_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f16_neon_packed_fused,
mr: 16,
nr: 4,
depth_multiple: 1,
};
#[cfg(all(feature = "half", target_arch = "aarch64"))]
const DISP_BF16_NEON: Dispatched<bf16> = Dispatched {
run: gemm_bf16_neon,
run_packed: gemm_bf16_neon_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_bf16_neon_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_bf16_neon_packed_fused,
mr: 16,
nr: 4,
depth_multiple: 1,
};
#[cfg(all(feature = "half", target_arch = "wasm32", target_feature = "simd128"))]
const DISP_F16_SIMD128: Dispatched<f16> = Dispatched {
run: gemm_f16_simd128,
run_packed: gemm_f16_simd128_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f16_simd128_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f16_simd128_packed_fused,
mr: 8,
nr: 4,
depth_multiple: 1,
};
#[cfg(all(feature = "half", target_arch = "wasm32", target_feature = "simd128"))]
const DISP_BF16_SIMD128: Dispatched<bf16> = Dispatched {
run: gemm_bf16_simd128,
run_packed: gemm_bf16_simd128_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_bf16_simd128_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_bf16_simd128_packed_fused,
mr: 8,
nr: 4,
depth_multiple: 1,
};
#[cfg(feature = "half")]
fn select_f16() -> Dispatched<f16> {
match forced_isa() {
ForcedIsa::Scalar => return DISP_F16_SCALAR,
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
ForcedIsa::Fma => {
assert!(
x86_isa_detected!("avx2") && x86_isa_detected!("fma") && x86_isa_detected!("f16c"),
"GEMMKIT_REQUIRE_ISA=fma for f16, but this CPU does not report avx2+fma+f16c"
);
return DISP_F16_FMA;
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
ForcedIsa::Avx512F | ForcedIsa::Avx512Vnni | ForcedIsa::Avx512Bf16 => {
assert!(
x86_isa_detected!("avx512f"),
"GEMMKIT_REQUIRE_ISA=avx512f, but this CPU/emulator does not report avx512f"
);
return DISP_F16_AVX512F;
}
#[cfg(not(any(target_arch = "x86", target_arch = "x86_64")))]
ForcedIsa::Fma | ForcedIsa::Avx512F | ForcedIsa::Avx512Vnni | ForcedIsa::Avx512Bf16 => {
panic!("GEMMKIT_REQUIRE_ISA: requested SIMD ISA is unavailable on this target")
}
#[cfg(target_arch = "aarch64")]
ForcedIsa::Neon => return DISP_F16_NEON,
#[cfg(not(target_arch = "aarch64"))]
ForcedIsa::Neon => panic!("GEMMKIT_REQUIRE_ISA=neon, but this target is not aarch64"),
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
ForcedIsa::Simd128 => return DISP_F16_SIMD128,
#[cfg(not(all(target_arch = "wasm32", target_feature = "simd128")))]
ForcedIsa::Simd128 => panic!(
"GEMMKIT_REQUIRE_ISA=simd128, but this build is not wasm32 with -C target-feature=+simd128"
),
ForcedIsa::Auto => {}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if x86_isa_detected!("avx512f") {
return DISP_F16_AVX512F;
}
if x86_isa_detected!("avx2") && x86_isa_detected!("fma") && x86_isa_detected!("f16c") {
return DISP_F16_FMA;
}
}
#[cfg(target_arch = "aarch64")]
{
DISP_F16_NEON
}
#[cfg(target_arch = "wasm32")]
{
#[cfg(target_feature = "simd128")]
{
DISP_F16_SIMD128
}
#[cfg(not(target_feature = "simd128"))]
{
DISP_F16_SCALAR
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "wasm32")))]
{
DISP_F16_SCALAR
}
}
#[cfg(feature = "half")]
fn select_bf16() -> Dispatched<bf16> {
match forced_isa() {
ForcedIsa::Scalar => return DISP_BF16_SCALAR,
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
ForcedIsa::Fma => {
assert!(
x86_isa_detected!("avx2") && x86_isa_detected!("fma"),
"GEMMKIT_REQUIRE_ISA=fma, but this CPU/emulator does not report avx2+fma"
);
return DISP_BF16_FMA;
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
ForcedIsa::Avx512F | ForcedIsa::Avx512Vnni => {
assert!(
x86_isa_detected!("avx512f"),
"GEMMKIT_REQUIRE_ISA=avx512f, but this CPU/emulator does not report avx512f"
);
return DISP_BF16_AVX512F;
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
ForcedIsa::Avx512Bf16 => {
assert!(
x86_isa_detected!("avx512bf16") && x86_isa_detected!("avx512f"),
"GEMMKIT_REQUIRE_ISA=avx512bf16, but this CPU/emulator does not report avx512f+bf16"
);
return DISP_BF16_AVX512BF16;
}
#[cfg(not(any(target_arch = "x86", target_arch = "x86_64")))]
ForcedIsa::Fma | ForcedIsa::Avx512F | ForcedIsa::Avx512Vnni | ForcedIsa::Avx512Bf16 => {
panic!("GEMMKIT_REQUIRE_ISA: requested SIMD ISA is unavailable on this target")
}
#[cfg(target_arch = "aarch64")]
ForcedIsa::Neon => return DISP_BF16_NEON,
#[cfg(not(target_arch = "aarch64"))]
ForcedIsa::Neon => panic!("GEMMKIT_REQUIRE_ISA=neon, but this target is not aarch64"),
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
ForcedIsa::Simd128 => return DISP_BF16_SIMD128,
#[cfg(not(all(target_arch = "wasm32", target_feature = "simd128")))]
ForcedIsa::Simd128 => panic!(
"GEMMKIT_REQUIRE_ISA=simd128, but this build is not wasm32 with -C target-feature=+simd128"
),
ForcedIsa::Auto => {}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if x86_isa_detected!("avx512bf16") && x86_isa_detected!("avx512f") {
return DISP_BF16_AVX512BF16;
}
if x86_isa_detected!("avx512f") {
return DISP_BF16_AVX512F;
}
if x86_isa_detected!("avx2") && x86_isa_detected!("fma") {
return DISP_BF16_FMA;
}
}
#[cfg(target_arch = "aarch64")]
{
DISP_BF16_NEON
}
#[cfg(target_arch = "wasm32")]
{
#[cfg(target_feature = "simd128")]
{
DISP_BF16_SIMD128
}
#[cfg(not(target_feature = "simd128"))]
{
DISP_BF16_SCALAR
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "wasm32")))]
{
DISP_BF16_SCALAR
}
}
memoized_select!(
GEMM_F16,
dispatched_f16,
Dispatched<f16>,
select_f16,
"The memoized dispatch descriptor for `f16` (selection runs once).",
"half"
);
memoized_select!(
GEMM_BF16,
dispatched_bf16,
Dispatched<bf16>,
select_bf16,
"The memoized dispatch descriptor for `bf16` (selection runs once).",
"half"
);
#[cfg(feature = "half")]
impl GemmScalar for f16 {
const OUT_IS_ACC: bool = false;
#[inline]
unsafe fn scale_c(beta: f16, c: *mut f16, m: usize, n: usize, rsc: isize, csc: isize) {
unsafe { scale_c_narrow(beta, c, m, n, rsc, csc) }
}
#[inline]
unsafe fn pack_rhs_full(
dst: *mut f16,
b: *const f16,
rsb: isize,
csb: isize,
k: usize,
n: usize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe { driver::pack_rhs_full::<MixedGemm<f16>>(dst, b, rsb, csb, k, n, kc, nc, nr) }
}
#[inline]
unsafe fn dispatch(task: Task<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { (dispatched_f16().run)(task, par, ws) }
}
#[inline]
unsafe fn dispatch_packed(req: PackedConsume<f16>, par: Parallelism, ws: &mut Workspace) {
unsafe { (dispatched_f16().run_packed)(req, par, ws) }
}
#[cfg(feature = "epilogue")]
#[inline]
unsafe fn dispatch_fused(
t: Task<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { (dispatched_f16().run_fused)(t, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
#[inline]
unsafe fn dispatch_packed_fused(
req: PackedConsume<f16>,
epi: FusedEpi<f16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { (dispatched_f16().run_packed_fused)(req, epi, par, ws) }
}
#[inline]
fn rhs_tile() -> (usize, usize) {
let d = dispatched_f16();
(d.mr, d.nr)
}
}
#[cfg(feature = "half")]
impl GemmScalar for bf16 {
const OUT_IS_ACC: bool = false;
#[inline]
unsafe fn scale_c(beta: bf16, c: *mut bf16, m: usize, n: usize, rsc: isize, csc: isize) {
unsafe { scale_c_narrow(beta, c, m, n, rsc, csc) }
}
#[inline]
unsafe fn pack_rhs_full(
dst: *mut bf16,
b: *const bf16,
rsb: isize,
csb: isize,
k: usize,
n: usize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
if dispatched_bf16().depth_multiple > 1 {
driver::pack_rhs_full::<Bf16DotGemm>(dst, b, rsb, csb, k, n, kc, nc, nr);
return;
}
driver::pack_rhs_full::<MixedGemm<bf16>>(dst, b, rsb, csb, k, n, kc, nc, nr);
}
}
#[inline]
unsafe fn dispatch(task: Task<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { (dispatched_bf16().run)(task, par, ws) }
}
#[inline]
unsafe fn dispatch_packed(req: PackedConsume<bf16>, par: Parallelism, ws: &mut Workspace) {
unsafe { (dispatched_bf16().run_packed)(req, par, ws) }
}
#[cfg(feature = "epilogue")]
#[inline]
unsafe fn dispatch_fused(
t: Task<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { (dispatched_bf16().run_fused)(t, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
#[inline]
unsafe fn dispatch_packed_fused(
req: PackedConsume<bf16>,
epi: FusedEpi<bf16>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { (dispatched_bf16().run_packed_fused)(req, epi, par, ws) }
}
#[inline]
fn rhs_tile() -> (usize, usize) {
let d = dispatched_bf16();
(d.mr, d.nr)
}
#[inline]
fn rhs_depth_multiple() -> usize {
dispatched_bf16().depth_multiple
}
}
#[cfg(all(feature = "half", feature = "epilogue"))]
impl FusedScalar for f16 {
#[inline]
fn finite(self) -> bool {
self.widen().is_finite()
}
#[inline]
unsafe fn fused_degenerate(t: &Task<f16>, epi: &FusedEpi<f16>) {
unsafe { fused_degenerate_mixed::<f16>(t, epi) }
}
}
#[cfg(all(feature = "half", feature = "epilogue"))]
impl FusedScalar for bf16 {
#[inline]
fn finite(self) -> bool {
self.widen().is_finite()
}
#[inline]
unsafe fn fused_degenerate(t: &Task<bf16>, epi: &FusedEpi<bf16>) {
unsafe { fused_degenerate_mixed::<bf16>(t, epi) }
}
}