#[cfg(feature = "std")]
use std::sync::OnceLock;
use super::isa::{ForcedIsa, forced_isa};
use super::{
GemmScalar, PackedConsume, Task, orient_transpose, scale_c_float, small_mn_eligible,
small_mn_pack_eligible,
};
use crate::driver;
use crate::kernel::FloatGemm;
use crate::kernel::epilogue::{Epilogue, Identity};
#[cfg(feature = "epilogue")]
use crate::kernel::epilogue::{FusedEpi, MapEpi};
use crate::parallel::Parallelism;
use crate::scalar::Float;
#[cfg(target_arch = "aarch64")]
use crate::simd::Neon;
#[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::{ScalarTok, SimdOps};
use crate::special::{gemv, small_k, small_mn};
use crate::tuning;
use crate::workspace::Workspace;
#[inline]
unsafe fn run_typed_epi<T, S, E, const MR_REG: usize, const NR: usize>(
simd: S,
mut t: Task<T>,
mut epi: E,
par: Parallelism,
ws: &mut Workspace,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
E: Epilogue<FloatGemm<T>>,
{
unsafe {
if (t.n == 1 || t.m == 1) && core::cmp::min(t.m, t.n) <= tuning::gemv_threshold() {
gemv::run_typed_epi::<T, S, E>(
simd, t.m, t.k, t.n, par, t.alpha, t.a, t.rsa, t.csa, t.b, t.rsb, t.csb, t.beta,
t.c, t.rsc, t.csc, &epi,
);
return;
}
if orient_transpose(&mut t) {
epi.on_orient_swap();
}
if small_mn_eligible(&t) || small_mn_pack_eligible(&t) {
small_mn::run_epi::<T, S, E>(
simd, t.m, t.k, t.n, par, ws, t.alpha, t.a, t.rsa, t.csa, t.b, t.rsb, t.csb,
t.beta, t.c, t.rsc, t.csc, &epi,
);
return;
}
if t.k <= tuning::small_k_threshold() {
small_k::run_epi::<FloatGemm<T>, S, E, MR_REG, NR>(
simd, t.m, t.k, t.n, t.alpha, t.a, t.rsa, t.csa, t.b, t.rsb, t.csb, t.beta, t.c,
t.rsc, t.csc, &epi, par, ws,
);
return;
}
driver::run_epilogue::<FloatGemm<T>, S, E, MR_REG, NR>(
simd, t.m, t.k, t.n, t.alpha, t.a, t.rsa, t.csa, t.b, t.rsb, t.csb, t.beta, t.c, t.rsc,
t.csc, &epi, par, ws,
);
}
}
#[inline]
unsafe fn run_typed<T, S, const MR_REG: usize, const NR: usize>(
simd: S,
t: Task<T>,
par: Parallelism,
ws: &mut Workspace,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
{
unsafe { run_typed_epi::<T, S, Identity, MR_REG, NR>(simd, t, Identity, par, ws) }
}
#[cfg(feature = "epilogue")]
#[inline]
unsafe fn run_typed_fused<T, S, const MR_REG: usize, const NR: usize>(
simd: S,
t: Task<T>,
epi: FusedEpi<T>,
par: Parallelism,
ws: &mut Workspace,
) where
T: Float<Acc = T> + PartialOrd,
S: SimdOps<T>,
{
unsafe { run_typed_epi::<T, S, FusedEpi<T>, MR_REG, NR>(simd, t, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
#[inline]
unsafe fn run_typed_map<'u, T, S, const MR_REG: usize, const NR: usize>(
simd: S,
t: Task<T>,
epi: MapEpi<'u, T>,
par: Parallelism,
ws: &mut Workspace,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
{
unsafe { run_typed_epi::<T, S, MapEpi<'u, T>, MR_REG, NR>(simd, t, epi, par, ws) }
}
#[inline]
unsafe fn run_packed_typed<T, S, const MR_REG: usize, const NR: usize>(
simd: S,
req: PackedConsume<T>,
par: Parallelism,
ws: &mut Workspace,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
{
unsafe {
debug_assert_eq!(NR, req.nr, "prepacked RHS panel width != kernel NR");
driver::run_packed_rhs::<FloatGemm<T>, S, MR_REG, NR>(
simd, req.m, req.k, req.n, req.alpha, req.a, req.rsa, req.csa, req.packed, req.kc,
req.nc, req.beta, req.c, req.rsc, req.csc, par, ws,
);
}
}
#[cfg(feature = "epilogue")]
#[inline]
unsafe fn run_typed_packed_fused<T, S, const MR_REG: usize, const NR: usize>(
simd: S,
req: PackedConsume<T>,
epi: FusedEpi<T>,
par: Parallelism,
ws: &mut Workspace,
) where
T: Float<Acc = T> + PartialOrd,
S: SimdOps<T>,
{
unsafe {
debug_assert_eq!(NR, req.nr, "prepacked RHS panel width != kernel NR");
driver::run_packed_rhs_epilogue::<FloatGemm<T>, S, FusedEpi<T>, MR_REG, NR>(
simd, req.m, req.k, req.n, req.alpha, req.a, req.rsa, req.csa, req.packed, req.kc,
req.nc, req.beta, req.c, req.rsc, req.csc, &epi, par, ws,
);
}
}
unsafe fn gemm_f32_scalar(t: Task<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f32, ScalarTok, 4, 4>(ScalarTok, t, par, ws) }
}
unsafe fn gemm_f64_scalar(t: Task<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f64, ScalarTok, 4, 4>(ScalarTok, t, par, ws) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn gemm_f32_fma(t: Task<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f32, Fma, 2, 6>(Fma, t, par, ws) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn gemm_f64_fma(t: Task<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f64, Fma, 2, 6>(Fma, t, par, ws) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn gemm_f32_avx512f(t: Task<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f32, Avx512F, 2, 12>(Avx512F, t, par, ws) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn gemm_f64_avx512f(t: Task<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f64, Avx512F, 2, 12>(Avx512F, t, par, ws) }
}
#[cfg(target_arch = "aarch64")]
unsafe fn gemm_f32_neon(t: Task<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f32, Neon, 4, 4>(Neon, t, par, ws) }
}
#[cfg(target_arch = "aarch64")]
unsafe fn gemm_f64_neon(t: Task<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f64, Neon, 4, 4>(Neon, t, par, ws) }
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
unsafe fn gemm_f32_simd128(t: Task<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f32, Simd128, 2, 4>(Simd128, t, par, ws) }
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
unsafe fn gemm_f64_simd128(t: Task<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_typed::<f64, Simd128, 2, 4>(Simd128, t, par, ws) }
}
unsafe fn gemm_f32_scalar_packed(r: PackedConsume<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f32, ScalarTok, 4, 4>(ScalarTok, r, par, ws) }
}
unsafe fn gemm_f64_scalar_packed(r: PackedConsume<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f64, ScalarTok, 4, 4>(ScalarTok, r, par, ws) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn gemm_f32_fma_packed(r: PackedConsume<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f32, Fma, 2, 6>(Fma, r, par, ws) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn gemm_f64_fma_packed(r: PackedConsume<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f64, Fma, 2, 6>(Fma, r, par, ws) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn gemm_f32_avx512f_packed(r: PackedConsume<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f32, Avx512F, 2, 12>(Avx512F, r, par, ws) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn gemm_f64_avx512f_packed(r: PackedConsume<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f64, Avx512F, 2, 12>(Avx512F, r, par, ws) }
}
#[cfg(target_arch = "aarch64")]
unsafe fn gemm_f32_neon_packed(r: PackedConsume<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f32, Neon, 4, 4>(Neon, r, par, ws) }
}
#[cfg(target_arch = "aarch64")]
unsafe fn gemm_f64_neon_packed(r: PackedConsume<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f64, Neon, 4, 4>(Neon, r, par, ws) }
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
unsafe fn gemm_f32_simd128_packed(r: PackedConsume<f32>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f32, Simd128, 2, 4>(Simd128, r, par, ws) }
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
unsafe fn gemm_f64_simd128_packed(r: PackedConsume<f64>, par: Parallelism, ws: &mut Workspace) {
unsafe { run_packed_typed::<f64, Simd128, 2, 4>(Simd128, r, par, ws) }
}
#[cfg(feature = "epilogue")]
unsafe fn gemm_f32_scalar_fused(
t: Task<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f32, ScalarTok, 4, 4>(ScalarTok, t, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
unsafe fn gemm_f64_scalar_fused(
t: Task<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f64, ScalarTok, 4, 4>(ScalarTok, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f32_fma_fused(
t: Task<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f32, Fma, 2, 6>(Fma, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f64_fma_fused(
t: Task<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f64, Fma, 2, 6>(Fma, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f32_avx512f_fused(
t: Task<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f32, Avx512F, 2, 12>(Avx512F, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f64_avx512f_fused(
t: Task<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f64, Avx512F, 2, 12>(Avx512F, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_f32_neon_fused(
t: Task<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f32, Neon, 4, 4>(Neon, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_f64_neon_fused(
t: Task<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f64, Neon, 4, 4>(Neon, t, epi, par, ws) }
}
#[cfg(all(
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_f32_simd128_fused(
t: Task<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f32, Simd128, 2, 4>(Simd128, t, epi, par, ws) }
}
#[cfg(all(
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_f64_simd128_fused(
t: Task<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_fused::<f64, Simd128, 2, 4>(Simd128, t, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
unsafe fn gemm_f32_scalar_packed_fused(
r: PackedConsume<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f32, ScalarTok, 4, 4>(ScalarTok, r, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
unsafe fn gemm_f64_scalar_packed_fused(
r: PackedConsume<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f64, ScalarTok, 4, 4>(ScalarTok, r, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f32_fma_packed_fused(
r: PackedConsume<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f32, Fma, 2, 6>(Fma, r, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f64_fma_packed_fused(
r: PackedConsume<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f64, Fma, 2, 6>(Fma, r, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f32_avx512f_packed_fused(
r: PackedConsume<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f32, Avx512F, 2, 12>(Avx512F, r, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f64_avx512f_packed_fused(
r: PackedConsume<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f64, Avx512F, 2, 12>(Avx512F, r, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_f32_neon_packed_fused(
r: PackedConsume<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f32, Neon, 4, 4>(Neon, r, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_f64_neon_packed_fused(
r: PackedConsume<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f64, Neon, 4, 4>(Neon, r, epi, par, ws) }
}
#[cfg(all(
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_f32_simd128_packed_fused(
r: PackedConsume<f32>,
epi: FusedEpi<f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f32, Simd128, 2, 4>(Simd128, r, epi, par, ws) }
}
#[cfg(all(
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_f64_simd128_packed_fused(
r: PackedConsume<f64>,
epi: FusedEpi<f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_packed_fused::<f64, Simd128, 2, 4>(Simd128, r, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
pub trait FusedScalar: GemmScalar + sealed::Sealed {
#[doc(hidden)]
fn finite(self) -> bool;
#[doc(hidden)]
unsafe fn fused_degenerate(t: &Task<Self>, epi: &FusedEpi<Self>);
}
#[cfg(feature = "epilogue")]
mod sealed {
pub trait Sealed {}
impl Sealed for f32 {}
impl Sealed for f64 {}
#[cfg(feature = "half")]
impl Sealed for half::f16 {}
#[cfg(feature = "half")]
impl Sealed for half::bf16 {}
}
#[cfg(feature = "epilogue")]
impl FusedScalar for f32 {
#[inline]
fn finite(self) -> bool {
self.is_finite()
}
#[inline]
unsafe fn fused_degenerate(t: &Task<f32>, epi: &FusedEpi<f32>) {
unsafe { fused_degenerate_float::<f32>(t, epi) }
}
}
#[cfg(feature = "epilogue")]
impl FusedScalar for f64 {
#[inline]
fn finite(self) -> bool {
self.is_finite()
}
#[inline]
unsafe fn fused_degenerate(t: &Task<f64>, epi: &FusedEpi<f64>) {
unsafe { fused_degenerate_float::<f64>(t, epi) }
}
}
#[cfg(feature = "epilogue")]
pub(crate) unsafe fn execute_fused<T: FusedScalar>(
task: Task<T>,
epi: FusedEpi<T>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
if task.m == 0 || task.n == 0 {
return;
}
if task.k == 0 || task.alpha == T::ZERO {
T::fused_degenerate(&task, &epi);
return;
}
T::dispatch_fused(task, epi, par, ws);
}
}
#[cfg(feature = "epilogue")]
pub(crate) unsafe fn execute_packed_fused<T: FusedScalar>(
req: PackedConsume<T>,
epi: FusedEpi<T>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
if req.m == 0 || req.n == 0 {
return;
}
if req.k == 0 || req.alpha == T::ZERO {
let task = Task {
m: req.m,
k: req.k,
n: req.n,
alpha: req.alpha,
a: req.a,
rsa: req.rsa,
csa: req.csa,
b: core::ptr::null(),
rsb: 0,
csb: 0,
beta: req.beta,
c: req.c,
rsc: req.rsc,
csc: req.csc,
};
T::fused_degenerate(&task, &epi);
return;
}
T::dispatch_packed_fused(req, epi, par, ws);
}
}
#[cfg(feature = "epilogue")]
pub(super) unsafe fn fused_degenerate_float<T: Float<Acc = T> + PartialOrd>(
t: &Task<T>,
epi: &FusedEpi<T>,
) {
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 = if t.beta == T::ZERO {
T::ZERO
} else if t.beta == T::ONE {
*p
} else {
t.beta * *p
};
*p = epi.apply(base, i, j);
}
}
}
}
type GemmFn<T> = unsafe fn(Task<T>, Parallelism, &mut Workspace);
type PackedFn<T> = unsafe fn(PackedConsume<T>, Parallelism, &mut Workspace);
#[cfg(feature = "epilogue")]
pub(super) type FusedFn<T> = unsafe fn(Task<T>, FusedEpi<T>, Parallelism, &mut Workspace);
#[cfg(feature = "epilogue")]
pub(super) type PackedFusedFn<T> =
unsafe fn(PackedConsume<T>, FusedEpi<T>, Parallelism, &mut Workspace);
#[derive(Copy, Clone)]
pub(super) struct Dispatched<T> {
pub(super) run: GemmFn<T>,
pub(super) run_packed: PackedFn<T>,
#[cfg(feature = "epilogue")]
pub(super) run_fused: FusedFn<T>,
#[cfg(feature = "epilogue")]
pub(super) run_packed_fused: PackedFusedFn<T>,
pub(super) mr: usize,
pub(super) nr: usize,
#[cfg_attr(not(feature = "half"), allow(dead_code))]
pub(super) depth_multiple: usize,
}
const DISP_F32_SCALAR: Dispatched<f32> = Dispatched {
run: gemm_f32_scalar,
run_packed: gemm_f32_scalar_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f32_scalar_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f32_scalar_packed_fused,
mr: 4,
nr: 4,
depth_multiple: 1,
};
const DISP_F64_SCALAR: Dispatched<f64> = Dispatched {
run: gemm_f64_scalar,
run_packed: gemm_f64_scalar_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f64_scalar_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f64_scalar_packed_fused,
mr: 4,
nr: 4,
depth_multiple: 1,
};
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
const DISP_F32_FMA: Dispatched<f32> = Dispatched {
run: gemm_f32_fma,
run_packed: gemm_f32_fma_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f32_fma_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f32_fma_packed_fused,
mr: 16,
nr: 6,
depth_multiple: 1,
};
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
const DISP_F64_FMA: Dispatched<f64> = Dispatched {
run: gemm_f64_fma,
run_packed: gemm_f64_fma_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f64_fma_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f64_fma_packed_fused,
mr: 8,
nr: 6,
depth_multiple: 1,
};
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
const DISP_F32_AVX512F: Dispatched<f32> = Dispatched {
run: gemm_f32_avx512f,
run_packed: gemm_f32_avx512f_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f32_avx512f_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f32_avx512f_packed_fused,
mr: 32,
nr: 12,
depth_multiple: 1,
};
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
const DISP_F64_AVX512F: Dispatched<f64> = Dispatched {
run: gemm_f64_avx512f,
run_packed: gemm_f64_avx512f_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f64_avx512f_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f64_avx512f_packed_fused,
mr: 16,
nr: 12,
depth_multiple: 1,
};
#[cfg(target_arch = "aarch64")]
const DISP_F32_NEON: Dispatched<f32> = Dispatched {
run: gemm_f32_neon,
run_packed: gemm_f32_neon_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f32_neon_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f32_neon_packed_fused,
mr: 16,
nr: 4,
depth_multiple: 1,
};
#[cfg(target_arch = "aarch64")]
const DISP_F64_NEON: Dispatched<f64> = Dispatched {
run: gemm_f64_neon,
run_packed: gemm_f64_neon_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f64_neon_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f64_neon_packed_fused,
mr: 8,
nr: 4,
depth_multiple: 1,
};
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
const DISP_F32_SIMD128: Dispatched<f32> = Dispatched {
run: gemm_f32_simd128,
run_packed: gemm_f32_simd128_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f32_simd128_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f32_simd128_packed_fused,
mr: 8,
nr: 4,
depth_multiple: 1,
};
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
const DISP_F64_SIMD128: Dispatched<f64> = Dispatched {
run: gemm_f64_simd128,
run_packed: gemm_f64_simd128_packed,
#[cfg(feature = "epilogue")]
run_fused: gemm_f64_simd128_fused,
#[cfg(feature = "epilogue")]
run_packed_fused: gemm_f64_simd128_packed_fused,
mr: 4,
nr: 4,
depth_multiple: 1,
};
fn select_f32() -> Dispatched<f32> {
match forced_isa() {
ForcedIsa::Scalar => return DISP_F32_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_F32_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_F32_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_F32_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_F32_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_F32_AVX512F;
}
if x86_isa_detected!("avx2") && x86_isa_detected!("fma") {
return DISP_F32_FMA;
}
}
#[cfg(target_arch = "aarch64")]
{
DISP_F32_NEON
}
#[cfg(target_arch = "wasm32")]
{
#[cfg(target_feature = "simd128")]
{
DISP_F32_SIMD128
}
#[cfg(not(target_feature = "simd128"))]
{
DISP_F32_SCALAR
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "wasm32")))]
{
DISP_F32_SCALAR
}
}
fn select_f64() -> Dispatched<f64> {
match forced_isa() {
ForcedIsa::Scalar => return DISP_F64_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_F64_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_F64_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_F64_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_F64_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_F64_AVX512F;
}
if x86_isa_detected!("avx2") && x86_isa_detected!("fma") {
return DISP_F64_FMA;
}
}
#[cfg(target_arch = "aarch64")]
{
DISP_F64_NEON
}
#[cfg(target_arch = "wasm32")]
{
#[cfg(target_feature = "simd128")]
{
DISP_F64_SIMD128
}
#[cfg(not(target_feature = "simd128"))]
{
DISP_F64_SCALAR
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "wasm32")))]
{
DISP_F64_SCALAR
}
}
memoized_select!(
GEMM_F32,
dispatched_f32,
Dispatched<f32>,
select_f32,
"The memoized dispatch descriptor for `f32` (selection runs once)."
);
memoized_select!(
GEMM_F64,
dispatched_f64,
Dispatched<f64>,
select_f64,
"The memoized dispatch descriptor for `f64` (selection runs once)."
);
macro_rules! float_gemm_scalar {
($t:ty, $disp:ident) => {
impl GemmScalar for $t {
const OUT_IS_ACC: bool = true;
#[inline]
unsafe fn scale_c(beta: $t, c: *mut $t, m: usize, n: usize, rsc: isize, csc: isize) {
unsafe { scale_c_float(beta, c, m, n, rsc, csc) }
}
#[inline]
unsafe fn pack_rhs_full(
dst: *mut $t,
b: *const $t,
rsb: isize,
csb: isize,
k: usize,
n: usize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe {
driver::pack_rhs_full::<FloatGemm<$t>>(dst, b, rsb, csb, k, n, kc, nc, nr)
}
}
#[inline]
unsafe fn dispatch(task: Task<$t>, par: Parallelism, ws: &mut Workspace) {
unsafe { ($disp().run)(task, par, ws) }
}
#[inline]
unsafe fn dispatch_packed(
req: PackedConsume<$t>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { ($disp().run_packed)(req, par, ws) }
}
#[cfg(feature = "epilogue")]
#[inline]
unsafe fn dispatch_fused(
t: Task<$t>,
epi: FusedEpi<$t>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { ($disp().run_fused)(t, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
#[inline]
unsafe fn dispatch_packed_fused(
req: PackedConsume<$t>,
epi: FusedEpi<$t>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { ($disp().run_packed_fused)(req, epi, par, ws) }
}
#[inline]
fn rhs_tile() -> (usize, usize) {
let d = $disp();
(d.mr, d.nr)
}
}
};
}
float_gemm_scalar!(f32, dispatched_f32);
float_gemm_scalar!(f64, dispatched_f64);
#[cfg(feature = "epilogue")]
type MapFn<T> = for<'u> unsafe fn(Task<T>, MapEpi<'u, T>, Parallelism, &mut Workspace);
#[cfg(feature = "epilogue")]
unsafe fn gemm_f32_scalar_map(
t: Task<f32>,
epi: MapEpi<'_, f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f32, ScalarTok, 4, 4>(ScalarTok, t, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
unsafe fn gemm_f64_scalar_map(
t: Task<f64>,
epi: MapEpi<'_, f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f64, ScalarTok, 4, 4>(ScalarTok, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f32_fma_map(
t: Task<f32>,
epi: MapEpi<'_, f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f32, Fma, 2, 6>(Fma, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f64_fma_map(
t: Task<f64>,
epi: MapEpi<'_, f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f64, Fma, 2, 6>(Fma, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f32_avx512f_map(
t: Task<f32>,
epi: MapEpi<'_, f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f32, Avx512F, 2, 12>(Avx512F, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", any(target_arch = "x86", target_arch = "x86_64")))]
unsafe fn gemm_f64_avx512f_map(
t: Task<f64>,
epi: MapEpi<'_, f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f64, Avx512F, 2, 12>(Avx512F, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_f32_neon_map(
t: Task<f32>,
epi: MapEpi<'_, f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f32, Neon, 4, 4>(Neon, t, epi, par, ws) }
}
#[cfg(all(feature = "epilogue", target_arch = "aarch64"))]
unsafe fn gemm_f64_neon_map(
t: Task<f64>,
epi: MapEpi<'_, f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f64, Neon, 4, 4>(Neon, t, epi, par, ws) }
}
#[cfg(all(
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_f32_simd128_map(
t: Task<f32>,
epi: MapEpi<'_, f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f32, Simd128, 2, 4>(Simd128, t, epi, par, ws) }
}
#[cfg(all(
feature = "epilogue",
target_arch = "wasm32",
target_feature = "simd128"
))]
unsafe fn gemm_f64_simd128_map(
t: Task<f64>,
epi: MapEpi<'_, f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { run_typed_map::<f64, Simd128, 2, 4>(Simd128, t, epi, par, ws) }
}
#[cfg(feature = "epilogue")]
fn select_map_f32() -> MapFn<f32> {
match forced_isa() {
ForcedIsa::Scalar => return gemm_f32_scalar_map,
#[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 gemm_f32_fma_map;
}
#[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 gemm_f32_avx512f_map;
}
#[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 gemm_f32_neon_map,
#[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 gemm_f32_simd128_map,
#[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 gemm_f32_avx512f_map;
}
if x86_isa_detected!("avx2") && x86_isa_detected!("fma") {
return gemm_f32_fma_map;
}
}
#[cfg(target_arch = "aarch64")]
{
gemm_f32_neon_map
}
#[cfg(target_arch = "wasm32")]
{
#[cfg(target_feature = "simd128")]
{
gemm_f32_simd128_map
}
#[cfg(not(target_feature = "simd128"))]
{
gemm_f32_scalar_map
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "wasm32")))]
{
gemm_f32_scalar_map
}
}
#[cfg(feature = "epilogue")]
fn select_map_f64() -> MapFn<f64> {
match forced_isa() {
ForcedIsa::Scalar => return gemm_f64_scalar_map,
#[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 gemm_f64_fma_map;
}
#[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 gemm_f64_avx512f_map;
}
#[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 gemm_f64_neon_map,
#[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 gemm_f64_simd128_map,
#[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 gemm_f64_avx512f_map;
}
if x86_isa_detected!("avx2") && x86_isa_detected!("fma") {
return gemm_f64_fma_map;
}
}
#[cfg(target_arch = "aarch64")]
{
gemm_f64_neon_map
}
#[cfg(target_arch = "wasm32")]
{
#[cfg(target_feature = "simd128")]
{
gemm_f64_simd128_map
}
#[cfg(not(target_feature = "simd128"))]
{
gemm_f64_scalar_map
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "wasm32")))]
{
gemm_f64_scalar_map
}
}
memoized_select!(
MAP_F32,
map_dispatched_f32,
MapFn<f32>,
select_map_f32,
"The memoized `f32` map-epilogue dispatch entry (selection runs once).",
"epilogue"
);
memoized_select!(
MAP_F64,
map_dispatched_f64,
MapFn<f64>,
select_map_f64,
"The memoized `f64` map-epilogue dispatch entry (selection runs once).",
"epilogue"
);
#[cfg(feature = "epilogue")]
pub trait MapScalar: GemmScalar + sealed::Sealed {
#[doc(hidden)]
unsafe fn dispatch_map(
task: Task<Self>,
epi: MapEpi<'_, Self>,
par: Parallelism,
ws: &mut Workspace,
);
#[doc(hidden)]
unsafe fn map_degenerate(t: &Task<Self>, epi: &MapEpi<'_, Self>);
}
#[cfg(feature = "epilogue")]
impl MapScalar for f32 {
#[inline]
unsafe fn dispatch_map(
task: Task<f32>,
epi: MapEpi<'_, f32>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { (map_dispatched_f32())(task, epi, par, ws) }
}
#[inline]
unsafe fn map_degenerate(t: &Task<f32>, epi: &MapEpi<'_, f32>) {
unsafe { map_degenerate_float::<f32>(t, epi) }
}
}
#[cfg(feature = "epilogue")]
impl MapScalar for f64 {
#[inline]
unsafe fn dispatch_map(
task: Task<f64>,
epi: MapEpi<'_, f64>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe { (map_dispatched_f64())(task, epi, par, ws) }
}
#[inline]
unsafe fn map_degenerate(t: &Task<f64>, epi: &MapEpi<'_, f64>) {
unsafe { map_degenerate_float::<f64>(t, epi) }
}
}
#[cfg(feature = "epilogue")]
pub(crate) unsafe fn execute_map<T: MapScalar>(
task: Task<T>,
epi: MapEpi<'_, T>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
if task.m == 0 || task.n == 0 {
return;
}
if task.k == 0 || task.alpha == T::ZERO {
T::map_degenerate(&task, &epi);
return;
}
T::dispatch_map(task, epi, par, ws);
}
}
#[cfg(feature = "epilogue")]
pub(super) unsafe fn map_degenerate_float<T: Float<Acc = T>>(t: &Task<T>, epi: &MapEpi<'_, T>) {
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 = if t.beta == T::ZERO {
T::ZERO
} else if t.beta == T::ONE {
*p
} else {
t.beta * *p
};
*p = epi.apply(base, i, j);
}
}
}
}