#![cfg_attr(
not(feature = "std"),
allow(
clippy::assertions_on_constants,
clippy::nonminimal_bool,
clippy::eq_op
)
)]
#[macro_use]
mod isa;
#[cfg(feature = "complex")]
mod complex;
mod float;
#[cfg(feature = "int8")]
mod int;
#[cfg(feature = "half")]
mod mixed;
#[cfg(feature = "complex")]
pub use complex::ComplexScalar;
#[cfg(feature = "complex")]
pub(crate) use complex::execute_complex;
#[cfg(all(feature = "complex", feature = "epilogue"))]
pub(crate) use complex::execute_complex_fused;
#[cfg(feature = "epilogue")]
pub use float::{FusedScalar, MapScalar};
#[cfg(feature = "epilogue")]
pub(crate) use float::{execute_fused, execute_map, execute_packed_fused};
#[cfg(feature = "int8")]
pub(crate) use int::{
IntPackedConsume, IntTask, execute_int, execute_int_packed, i8_rhs_depth_multiple, i8_rhs_tile,
pack_rhs_full_i8,
};
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub(crate) use int::{RequantTask, execute_int_requant};
#[cfg(feature = "epilogue")]
use crate::kernel::epilogue::FusedEpi;
use crate::parallel::Parallelism;
use crate::scalar::{Float, Scalar};
use crate::tuning;
use crate::workspace::Workspace;
#[derive(Copy, Clone)]
pub struct Task<T> {
pub m: usize,
pub k: usize,
pub n: usize,
pub alpha: T,
pub a: *const T,
pub rsa: isize,
pub csa: isize,
pub b: *const T,
pub rsb: isize,
pub csb: isize,
pub beta: T,
pub c: *mut T,
pub rsc: isize,
pub csc: isize,
}
#[derive(Copy, Clone)]
pub struct GemmProblem<T> {
pub m: usize,
pub k: usize,
pub n: usize,
pub alpha: T,
pub a: *const T,
pub rsa: isize,
pub csa: isize,
pub b: *const T,
pub rsb: isize,
pub csb: isize,
pub beta: T,
pub c: *mut T,
pub rsc: isize,
pub csc: isize,
}
impl<T: Copy> GemmProblem<T> {
#[inline]
pub(crate) fn task(&self) -> Task<T> {
Task {
m: self.m,
k: self.k,
n: self.n,
alpha: self.alpha,
a: self.a,
rsa: self.rsa,
csa: self.csa,
b: self.b,
rsb: self.rsb,
csb: self.csb,
beta: self.beta,
c: self.c,
rsc: self.rsc,
csc: self.csc,
}
}
}
pub struct PackedConsume<T> {
pub m: usize,
pub k: usize,
pub n: usize,
pub alpha: T,
pub a: *const T,
pub rsa: isize,
pub csa: isize,
pub packed: *const T,
pub nr: usize,
pub kc: usize,
pub nc: usize,
pub beta: T,
pub c: *mut T,
pub rsc: isize,
pub csc: isize,
}
pub trait GemmScalar: Scalar {
const OUT_IS_ACC: bool;
#[doc(hidden)]
unsafe fn scale_c(beta: Self, c: *mut Self, m: usize, n: usize, rsc: isize, csc: isize);
#[doc(hidden)]
#[allow(clippy::too_many_arguments)]
unsafe fn pack_rhs_full(
dst: *mut Self,
b: *const Self,
rsb: isize,
csb: isize,
k: usize,
n: usize,
kc: usize,
nc: usize,
nr: usize,
);
#[doc(hidden)]
unsafe fn dispatch(task: Task<Self>, par: Parallelism, ws: &mut Workspace);
#[doc(hidden)]
unsafe fn dispatch_packed(req: PackedConsume<Self>, par: Parallelism, ws: &mut Workspace);
#[doc(hidden)]
fn rhs_tile() -> (usize, usize);
#[doc(hidden)]
fn rhs_depth_multiple() -> usize {
1
}
#[doc(hidden)]
#[cfg(feature = "epilogue")]
unsafe fn dispatch_fused(
task: Task<Self>,
epi: FusedEpi<Self>,
par: Parallelism,
ws: &mut Workspace,
);
#[doc(hidden)]
#[cfg(feature = "epilogue")]
unsafe fn dispatch_packed_fused(
req: PackedConsume<Self>,
epi: FusedEpi<Self>,
par: Parallelism,
ws: &mut Workspace,
);
}
pub(crate) unsafe fn execute<T: GemmScalar>(task: Task<T>, par: Parallelism, ws: &mut Workspace) {
unsafe {
if task.m == 0 || task.n == 0 {
return;
}
if task.k == 0 || task.alpha == T::ZERO {
T::scale_c(task.beta, task.c, task.m, task.n, task.rsc, task.csc);
return;
}
T::dispatch(task, par, ws);
}
}
pub(crate) unsafe fn execute_packed<T: GemmScalar>(
req: PackedConsume<T>,
par: Parallelism,
ws: &mut Workspace,
) {
unsafe {
if req.m == 0 || req.n == 0 {
return;
}
if req.k == 0 || req.alpha == T::ZERO {
T::scale_c(req.beta, req.c, req.m, req.n, req.rsc, req.csc);
return;
}
T::dispatch_packed(req, par, ws);
}
}
#[inline]
#[allow(clippy::too_many_arguments)]
fn orient_swap<L>(
m: &mut usize,
n: &mut usize,
a: &mut *const L,
rsa: &mut isize,
csa: &mut isize,
b: &mut *const L,
rsb: &mut isize,
csb: &mut isize,
rsc: &mut isize,
csc: &mut isize,
) -> bool {
if csc.unsigned_abs() < rsc.unsigned_abs() {
core::mem::swap(m, n);
core::mem::swap(a, b); core::mem::swap(rsa, csb); core::mem::swap(csa, rsb); core::mem::swap(rsc, csc);
true
} else {
false
}
}
#[inline]
fn orient_transpose<T>(t: &mut Task<T>) -> bool {
orient_swap(
&mut t.m, &mut t.n, &mut t.a, &mut t.rsa, &mut t.csa, &mut t.b, &mut t.rsb, &mut t.csb,
&mut t.rsc, &mut t.csc,
)
}
#[inline]
fn small_mn_eligible_dims(m: usize, n: usize, k: usize, csa: isize, rsb: isize) -> bool {
m <= tuning::small_mn_dim()
&& n <= tuning::small_mn_dim()
&& k > tuning::small_k_threshold()
&& csa == 1
&& rsb == 1
}
#[inline]
fn small_mn_eligible<T>(t: &Task<T>) -> bool {
small_mn_eligible_dims(t.m, t.n, t.k, t.csa, t.rsb)
}
#[inline]
fn small_mn_pack_eligible_dims(m: usize, n: usize, k: usize, csa: isize, rsb: isize) -> bool {
m <= tuning::small_mn_dim()
&& n <= tuning::small_mn_dim()
&& k > tuning::small_mn_pack_min_k()
&& !(csa == 1 && rsb == 1)
}
#[inline]
fn small_mn_pack_eligible<T>(t: &Task<T>) -> bool {
small_mn_pack_eligible_dims(t.m, t.n, t.k, t.csa, t.rsb)
}
unsafe fn scale_c_float<T: Float>(beta: T, c: *mut T, m: usize, n: usize, rsc: isize, csc: isize) {
unsafe {
for j in 0..n {
for i in 0..m {
let p = c.offset(i as isize * rsc + j as isize * csc);
if beta == T::ZERO {
*p = T::ZERO;
} else if beta != T::ONE {
*p = beta * *p;
}
}
}
}
}