use super::KernelFamily;
#[cfg(feature = "epilogue")]
use super::float::FloatGemm;
#[cfg(feature = "epilogue")]
use crate::parallel::Ptr;
#[cfg(feature = "epilogue")]
use crate::scalar::Float;
use crate::simd::{KernelSimd, SimdOps};
pub trait Epilogue<Fam: KernelFamily>: Copy + Send + Sync {
const IS_IDENTITY: bool = false;
const VECTOR: bool = false;
unsafe fn apply(&self, v: Fam::Acc, row: usize, col: usize) -> Fam::Out;
#[inline(always)]
unsafe fn apply_reg<S>(
&self,
_simd: S,
_v: <S as SimdOps<Fam::Acc>>::Reg,
_row: usize,
_col: usize,
) -> <S as SimdOps<Fam::Acc>>::Reg
where
S: KernelSimd<Fam::Lhs, Fam::Rhs, Fam::Acc, Fam::Out>,
{
unreachable!("apply_reg requires VECTOR = true")
}
const VECTOR_STORE: bool = false;
#[inline(always)]
unsafe fn apply_store<S>(
&self,
_simd: S,
_src: *const Fam::Acc,
_dst: *mut Fam::Out,
_row: usize,
_col: usize,
) where
S: KernelSimd<Fam::Lhs, Fam::Rhs, Fam::Acc, Fam::Out>,
{
unreachable!("apply_store requires VECTOR_STORE = true")
}
#[inline(always)]
fn on_orient_swap(&mut self) {}
}
#[derive(Copy, Clone, Default)]
pub struct Identity;
impl<Fam: KernelFamily> Epilogue<Fam> for Identity {
const IS_IDENTITY: bool = true;
#[inline(always)]
unsafe fn apply(&self, _: Fam::Acc, _: usize, _: usize) -> Fam::Out {
unreachable!("identity epilogue is never applied")
}
}
#[cfg(feature = "epilogue")]
#[derive(Copy, Clone)]
pub enum BiasDim {
PerRow,
PerCol,
}
#[cfg(feature = "epilogue")]
#[derive(Copy, Clone)]
pub(crate) enum BiasSpec<T> {
None,
Row(Ptr<T>),
Col(Ptr<T>),
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
#[derive(Copy, Clone)]
pub(crate) enum ScaleSpec {
Tensor(f32),
Row(Ptr<f32>),
Col(Ptr<f32>),
}
#[cfg(feature = "epilogue")]
#[derive(Copy, Clone)]
pub(crate) enum Act<T> {
None,
Relu,
LeakyRelu(T),
}
#[cfg(feature = "epilogue")]
#[derive(Copy, Clone)]
pub struct FusedEpi<T> {
pub(crate) bias: BiasSpec<T>,
pub(crate) act: Act<T>,
}
#[cfg(feature = "epilogue")]
impl<T> FusedEpi<T> {
#[inline(always)]
pub(crate) fn flip_bias(&mut self) {
self.bias = match &self.bias {
BiasSpec::None => BiasSpec::None,
BiasSpec::Row(p) => BiasSpec::Col(Ptr(p.0)),
BiasSpec::Col(p) => BiasSpec::Row(Ptr(p.0)),
};
}
}
#[cfg(feature = "epilogue")]
impl<T: Float<Acc = T> + PartialOrd> Epilogue<FloatGemm<T>> for FusedEpi<T> {
const VECTOR: bool = true;
#[inline(always)]
fn on_orient_swap(&mut self) {
self.flip_bias();
}
#[inline(always)]
unsafe fn apply(&self, v: T, r: usize, c: usize) -> T {
let v = match self.bias {
BiasSpec::None => v,
BiasSpec::Row(p) => v + unsafe { *p.0.add(r) },
BiasSpec::Col(p) => v + unsafe { *p.0.add(c) },
};
match self.act {
Act::None => v,
Act::Relu => {
if v > T::ZERO {
v
} else {
T::ZERO
}
}
Act::LeakyRelu(s) => {
let hi = if v > T::ZERO { v } else { T::ZERO };
let lo = if v < T::ZERO { v } else { T::ZERO };
hi + s * lo
}
}
}
#[inline(always)]
unsafe fn apply_reg<S>(&self, s: S, v: S::Reg, r: usize, c: usize) -> S::Reg
where
S: KernelSimd<T, T, T, T>,
{
unsafe {
let v = match self.bias {
BiasSpec::None => v,
BiasSpec::Row(p) => s.add(v, s.loadu(p.0.add(r))),
BiasSpec::Col(p) => s.add(v, s.splat(*p.0.add(c))),
};
match self.act {
Act::None => v,
Act::Relu => s.max(v, s.zero()),
Act::LeakyRelu(sl) => {
s.add(s.max(v, s.zero()), s.mul(s.splat(sl), s.min(v, s.zero())))
}
}
}
}
}
#[cfg(feature = "epilogue")]
pub struct MapEpi<'u, T> {
pub(crate) f: &'u (dyn Fn(T, usize, usize) -> T + Sync),
pub(crate) swapped: bool,
}
#[cfg(feature = "epilogue")]
impl<T> Copy for MapEpi<'_, T> {}
#[cfg(feature = "epilogue")]
impl<T> Clone for MapEpi<'_, T> {
fn clone(&self) -> Self {
*self
}
}
#[cfg(feature = "epilogue")]
impl<T: Float<Acc = T>> Epilogue<FloatGemm<T>> for MapEpi<'_, T> {
const VECTOR: bool = true;
#[inline(always)]
fn on_orient_swap(&mut self) {
self.swapped = true;
}
#[inline]
unsafe fn apply(&self, v: T, r: usize, c: usize) -> T {
if self.swapped {
(self.f)(v, c, r)
} else {
(self.f)(v, r, c)
}
}
#[inline]
unsafe fn apply_reg<S>(&self, s: S, v: S::Reg, r: usize, c: usize) -> S::Reg
where
S: KernelSimd<T, T, T, T>,
{
unsafe {
let lanes = <S as SimdOps<T>>::LANES;
debug_assert!(
lanes <= MAP_REG_LANES,
"map apply_reg buffer holds MAP_REG_LANES lanes"
);
let mut buf = [T::ZERO; MAP_REG_LANES];
s.storeu(buf.as_mut_ptr(), v);
for (l, slot) in buf.iter_mut().enumerate().take(lanes) {
*slot = self.apply(*slot, r + l, c);
}
s.loadu(buf.as_ptr())
}
}
}
#[cfg(feature = "epilogue")]
const MAP_REG_LANES: usize = 16;
#[cfg(all(feature = "half", feature = "epilogue"))]
impl<N, Fam> Epilogue<Fam> for FusedEpi<N>
where
N: crate::scalar::NarrowFloat,
Fam: KernelFamily<Lhs = N, Rhs = N, Acc = f32, Out = N>,
{
const VECTOR: bool = true;
#[inline(always)]
fn on_orient_swap(&mut self) {
self.flip_bias();
}
#[inline(always)]
unsafe fn apply(&self, v: f32, r: usize, c: usize) -> N {
let v = match self.bias {
BiasSpec::None => v,
BiasSpec::Row(p) => v + unsafe { (*p.0.add(r)).widen() },
BiasSpec::Col(p) => v + unsafe { (*p.0.add(c)).widen() },
};
let v = match self.act {
Act::None => v,
Act::Relu => {
if v > 0.0 {
v
} else {
0.0
}
}
Act::LeakyRelu(s) => {
let hi = if v > 0.0 { v } else { 0.0 };
let lo = if v < 0.0 { v } else { 0.0 };
hi + s.widen() * lo
}
};
N::narrow(v)
}
#[inline(always)]
unsafe fn apply_reg<S>(&self, s: S, v: S::Reg, r: usize, c: usize) -> S::Reg
where
S: KernelSimd<N, N, f32, N>,
{
unsafe {
let v = match self.bias {
BiasSpec::None => v,
BiasSpec::Row(p) => s.add(v, s.load_lhs(p.0.add(r))),
BiasSpec::Col(p) => s.add(v, s.splat((*p.0.add(c)).widen())),
};
match self.act {
Act::None => v,
Act::Relu => s.max(v, s.zero()),
Act::LeakyRelu(sl) => s.add(
s.max(v, s.zero()),
s.mul(s.splat(sl.widen()), s.min(v, s.zero())),
),
}
}
}
}
#[cfg(all(feature = "complex", feature = "epilogue"))]
impl<T, const CA: bool, const CB: bool> Epilogue<crate::kernel::ComplexGemm<T, CA, CB>>
for FusedEpi<T>
where
T: crate::scalar::ComplexFloat,
{
#[inline(always)]
fn on_orient_swap(&mut self) {
self.flip_bias();
}
#[inline(always)]
unsafe fn apply(&self, v: T, r: usize, c: usize) -> T {
let v = match self.bias {
BiasSpec::None => v,
BiasSpec::Row(p) => v + unsafe { *p.0.add(r) },
BiasSpec::Col(p) => v + unsafe { *p.0.add(c) },
};
match self.act {
Act::None => v,
Act::Relu | Act::LeakyRelu(_) => {
unreachable!("complex fused epilogue has no activation")
}
}
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub(crate) trait QuantOut: crate::scalar::Scalar {
const LO: i32;
const HI: i32;
fn from_clamped(q: i64) -> Self;
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
impl QuantOut for i8 {
const LO: i32 = -128;
const HI: i32 = 127;
#[inline(always)]
fn from_clamped(q: i64) -> Self {
q as i8
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
impl QuantOut for u8 {
const LO: i32 = 0;
const HI: i32 = 255;
#[inline(always)]
fn from_clamped(q: i64) -> Self {
q as u8
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
#[derive(Copy, Clone)]
pub(crate) struct KRequantize {
pub scale: ScaleSpec,
pub zp: i32,
pub bias: BiasSpec<i32>,
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
impl<O: QuantOut, Fam: KernelFamily<Acc = i32, Out = O>> Epilogue<Fam> for KRequantize {
const VECTOR: bool = false;
const VECTOR_STORE: bool = true;
#[inline(always)]
unsafe fn apply(&self, v: i32, r: usize, c: usize) -> O {
let b = match self.bias {
BiasSpec::None => 0,
BiasSpec::Row(p) => unsafe { *p.0.add(r) },
BiasSpec::Col(p) => unsafe { *p.0.add(c) },
};
let scale = match self.scale {
ScaleSpec::Tensor(s) => s,
ScaleSpec::Row(p) => unsafe { *p.0.add(r) },
ScaleSpec::Col(p) => unsafe { *p.0.add(c) },
};
let scaled = round_ne_f64(f64::from(v.wrapping_add(b)) * f64::from(scale));
let q = (scaled as i64).saturating_add(i64::from(self.zp));
O::from_clamped(q.clamp(i64::from(O::LO), i64::from(O::HI)))
}
#[inline(always)]
unsafe fn apply_store<S>(&self, simd: S, src: *const i32, dst: *mut O, row: usize, col: usize)
where
S: KernelSimd<Fam::Lhs, Fam::Rhs, i32, O>,
{
unsafe {
if let ScaleSpec::Row(_) = self.scale {
let lanes = <S as SimdOps<i32>>::LANES;
for l in 0..lanes {
*dst.add(l) = <Self as Epilogue<Fam>>::apply(self, *src.add(l), row + l, col);
}
return;
}
let v = simd.loadu(src);
let v = match self.bias {
BiasSpec::None => v,
BiasSpec::Row(p) => simd.add(v, simd.loadu(p.0.add(row))),
BiasSpec::Col(p) => simd.add(v, simd.splat(*p.0.add(col))),
};
let scale = match self.scale {
ScaleSpec::Tensor(s) => f64::from(s),
ScaleSpec::Col(p) => f64::from(*p.0.add(col)),
ScaleSpec::Row(_) => unreachable!("per-row scale takes the per-lane path"),
};
simd.requant_store(dst as *mut i8, v, scale, self.zp, O::LO, O::HI);
}
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
#[inline(always)]
pub(crate) fn round_ne_f64(x: f64) -> f64 {
const C: f64 = 4503599627370496.0; if x.is_nan() || x >= C || x <= -C {
x
} else if x >= 0.0 {
(x + C) - C
} else {
(x - C) + C
}
}