use crate::dispatch::{self, GemmScalar, Task};
#[cfg(feature = "epilogue")]
use crate::kernel::epilogue::{Act, BiasDim, BiasSpec, FusedEpi};
use crate::parallel::Parallelism;
#[cfg(feature = "epilogue")]
use crate::parallel::Ptr;
use crate::workspace::{self, Workspace};
mod batched;
#[cfg(feature = "complex")]
mod cplx;
#[cfg(feature = "epilogue")]
mod fused;
#[cfg(feature = "int8")]
mod int8;
#[cfg(feature = "epilogue")]
mod map;
mod packed;
pub use batched::{
BatchProblem, gemm_batched, gemm_batched_ptr_unchecked, gemm_batched_slice,
gemm_batched_unchecked, gemm_batched_unchecked_with, gemm_batched_with,
};
#[cfg(feature = "epilogue")]
pub use batched::{
gemm_batched_fused, gemm_batched_fused_unchecked, gemm_batched_fused_unchecked_with,
gemm_batched_fused_with,
};
#[cfg(feature = "complex")]
pub use cplx::{gemm_cplx, gemm_cplx_unchecked, gemm_cplx_unchecked_with, gemm_cplx_with};
#[cfg(all(feature = "complex", feature = "epilogue"))]
pub use cplx::{
gemm_cplx_fused, gemm_cplx_fused_unchecked, gemm_cplx_fused_unchecked_with,
gemm_cplx_fused_with,
};
#[cfg(feature = "epilogue")]
pub use fused::{
Activation, Bias, gemm_fused, gemm_fused_unchecked, gemm_fused_unchecked_with, gemm_fused_with,
};
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub use int8::{
RequantScale, Requantize, gemm_i8_requant, gemm_i8_requant_u8, gemm_i8_requant_u8_unchecked,
gemm_i8_requant_u8_unchecked_with, gemm_i8_requant_u8_with, gemm_i8_requant_unchecked,
gemm_i8_requant_unchecked_with, gemm_i8_requant_with,
};
#[cfg(feature = "int8")]
pub use int8::{gemm_i8, gemm_i8_unchecked, gemm_i8_unchecked_with, gemm_i8_with};
#[cfg(feature = "epilogue")]
pub use map::{gemm_map, gemm_map_unchecked, gemm_map_unchecked_with, gemm_map_with};
pub use packed::{
PackedLhs, PackedRhs, gemm_packed_a, gemm_packed_a_unchecked, gemm_packed_a_unchecked_with,
gemm_packed_a_with, gemm_packed_b, gemm_packed_b_unchecked, gemm_packed_b_unchecked_with,
gemm_packed_b_with, prepack_lhs, prepack_lhs_unchecked, prepack_rhs, prepack_rhs_unchecked,
};
#[cfg(feature = "int8")]
pub use packed::{
gemm_i8_packed_b, gemm_i8_packed_b_unchecked, gemm_i8_packed_b_unchecked_with,
gemm_i8_packed_b_with, prepack_rhs_i8, prepack_rhs_i8_unchecked,
};
#[cfg(feature = "epilogue")]
pub use packed::{
gemm_packed_a_fused, gemm_packed_a_fused_unchecked, gemm_packed_a_fused_unchecked_with,
gemm_packed_a_fused_with, gemm_packed_b_fused, gemm_packed_b_fused_unchecked,
gemm_packed_b_fused_unchecked_with, gemm_packed_b_fused_with,
};
#[derive(Copy, Clone)]
pub struct MatRef<'a, T> {
data: &'a [T],
rows: usize,
cols: usize,
rs: isize,
cs: isize,
}
pub struct MatMut<'a, T> {
data: &'a mut [T],
rows: usize,
cols: usize,
rs: isize,
cs: isize,
}
impl<'a, T> MatRef<'a, T> {
pub fn new(data: &'a [T], rows: usize, cols: usize, rs: isize, cs: isize) -> Self {
Self {
data,
rows,
cols,
rs,
cs,
}
}
pub fn from_row_major(data: &'a [T], rows: usize, cols: usize) -> Self {
Self::new(data, rows, cols, cols as isize, 1)
}
pub fn from_col_major(data: &'a [T], rows: usize, cols: usize) -> Self {
Self::new(data, rows, cols, 1, rows as isize)
}
pub fn rows(&self) -> usize {
self.rows
}
pub fn cols(&self) -> usize {
self.cols
}
}
impl<'a, T> MatMut<'a, T> {
pub fn new(data: &'a mut [T], rows: usize, cols: usize, rs: isize, cs: isize) -> Self {
Self {
data,
rows,
cols,
rs,
cs,
}
}
pub fn from_row_major(data: &'a mut [T], rows: usize, cols: usize) -> Self {
let cs = cols as isize;
Self::new(data, rows, cols, cs, 1)
}
pub fn from_col_major(data: &'a mut [T], rows: usize, cols: usize) -> Self {
let rs = rows as isize;
Self::new(data, rows, cols, 1, rs)
}
pub fn rows(&self) -> usize {
self.rows
}
pub fn cols(&self) -> usize {
self.cols
}
}
fn extent(rows: usize, cols: usize, rs: isize, cs: isize) -> Option<usize> {
if rows == 0 || cols == 0 {
return Some(0);
}
let mut lo: isize = 0;
let mut hi: isize = 0;
for &(dim, s) in &[(rows, rs), (cols, cs)] {
let e = isize::try_from(dim).ok()?.checked_sub(1)?.checked_mul(s)?;
if e < 0 {
lo = lo.checked_add(e)?;
} else {
hi = hi.checked_add(e)?;
}
}
if lo < 0 {
None } else {
(hi as usize).checked_add(1)
}
}
fn check_view<T>(data: &[T], rows: usize, cols: usize, rs: isize, cs: isize, name: &str) {
match extent(rows, cols, rs, cs) {
Some(need) if need <= data.len() => {}
Some(need) => panic!(
"gemmkit: {name} view of {rows}x{cols} (strides {rs},{cs}) needs {need} elements but slice has {}",
data.len()
),
None => panic!(
"gemmkit: {name} view has negative strides or is too large to address; use gemm_unchecked"
),
}
}
fn self_aliases(rows: usize, cols: usize, rs: isize, cs: isize) -> bool {
if rows == 0 || cols == 0 {
return false; }
let r = (rows > 1).then_some((rs.unsigned_abs(), rows));
let c = (cols > 1).then_some((cs.unsigned_abs(), cols));
match (r, c) {
(None, None) => false,
(Some((s, _)), None) | (None, Some((s, _))) => s == 0,
(Some(a), Some(b)) => {
let (sm, big) = if a.0 <= b.0 { (a, b.0) } else { (b, a.0) };
sm.0 == 0 || big < sm.0.saturating_mul(sm.1)
}
}
}
fn overlaps<T>(pa: *const T, na: usize, pb: *const T, nb: usize) -> bool {
let s = core::mem::size_of::<T>();
overlaps_bytes(pa as *const u8, na, s, pb as *const u8, nb, s)
}
fn validate_gemm_views<TI, TO>(a: &MatRef<'_, TI>, b: &MatRef<'_, TI>, c: &MatMut<'_, TO>) {
assert_eq!(
a.cols, b.rows,
"gemmkit: A.cols ({}) != B.rows ({})",
a.cols, b.rows
);
assert_eq!(
a.rows, c.rows,
"gemmkit: A.rows ({}) != C.rows ({})",
a.rows, c.rows
);
assert_eq!(
b.cols, c.cols,
"gemmkit: B.cols ({}) != C.cols ({})",
b.cols, c.cols
);
check_view(a.data, a.rows, a.cols, a.rs, a.cs, "A");
check_view(b.data, b.rows, b.cols, b.rs, b.cs, "B");
check_view(c.data, c.rows, c.cols, c.rs, c.cs, "C");
if self_aliases(c.rows, c.cols, c.rs, c.cs) {
panic!(
"gemmkit: C view aliases itself (strides {},{} map distinct elements to the same \
memory); C must address each (i,j) uniquely",
c.rs, c.cs
);
}
let cp = c.data.as_ptr() as *const u8;
let cl = c.data.len();
let si = core::mem::size_of::<TI>();
let so = core::mem::size_of::<TO>();
if overlaps_bytes(cp, cl, so, a.data.as_ptr() as *const u8, a.data.len(), si)
|| overlaps_bytes(cp, cl, so, b.data.as_ptr() as *const u8, b.data.len(), si)
{
panic!("gemmkit: C aliases A or B");
}
}
#[cfg(feature = "epilogue")]
fn validate_bias<T: Copy>(bias: &Option<Bias<'_, T>>, m: usize, n: usize, c: &MatMut<'_, T>) {
let _ = crate::adapter::lower_bias(*bias, m, n, c.data.as_ptr(), &[(c.data.len(), 1)]);
}
#[cfg(feature = "epilogue")]
fn to_fused_epi<T>(bias: Option<Bias<'_, T>>, act: Option<Activation<T>>) -> FusedEpi<T> {
let bias = match bias {
None => BiasSpec::None,
Some(Bias::PerRow(s)) => BiasSpec::Row(Ptr(s.as_ptr() as *mut T)),
Some(Bias::PerCol(s)) => BiasSpec::Col(Ptr(s.as_ptr() as *mut T)),
};
let act = match act {
None => Act::None,
Some(Activation::Relu) => Act::Relu,
Some(Activation::LeakyRelu(s)) => Act::LeakyRelu(s),
};
FusedEpi { bias, act }
}
#[cfg(feature = "epilogue")]
fn to_fused_epi_raw<T>(
bias: *const T,
bias_dim: BiasDim,
has_bias: bool,
act: Option<Activation<T>>,
) -> FusedEpi<T> {
let bias = if has_bias {
match bias_dim {
BiasDim::PerRow => BiasSpec::Row(Ptr(bias as *mut T)),
BiasDim::PerCol => BiasSpec::Col(Ptr(bias as *mut T)),
}
} else {
BiasSpec::None
};
let act = match act {
None => Act::None,
Some(Activation::Relu) => Act::Relu,
Some(Activation::LeakyRelu(s)) => Act::LeakyRelu(s),
};
FusedEpi { bias, act }
}
pub fn gemm<T: GemmScalar>(
alpha: T,
a: MatRef<'_, T>,
b: MatRef<'_, T>,
beta: T,
c: MatMut<'_, T>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| gemm_with(ws, alpha, a, b, beta, c, par));
}
pub fn gemm_with<T: GemmScalar>(
ws: &mut Workspace,
alpha: T,
a: MatRef<'_, T>,
b: MatRef<'_, T>,
beta: T,
c: MatMut<'_, T>,
par: Parallelism,
) {
validate_gemm_views(&a, &b, &c);
let m = a.rows;
let k = a.cols;
let n = b.cols;
unsafe {
dispatch::execute(
Task {
m,
k,
n,
alpha,
a: a.data.as_ptr(),
rsa: a.rs,
csa: a.cs,
b: b.data.as_ptr(),
rsb: b.rs,
csb: b.cs,
beta,
c: c.data.as_mut_ptr(),
rsc: c.rs,
csc: c.cs,
},
par,
ws,
);
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_unchecked<T: GemmScalar>(
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
b: *const T,
rsb: isize,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
dispatch::execute(
Task {
m,
k,
n,
alpha,
a,
rsa,
csa,
b,
rsb,
csb,
beta,
c,
rsc,
csc,
},
par,
ws,
);
});
}
}
fn overlaps_bytes(
pa: *const u8,
na: usize,
sa: usize,
pb: *const u8,
nb: usize,
sb: usize,
) -> bool {
let a0 = pa as usize;
let a1 = a0 + na * sa;
let b0 = pb as usize;
let b1 = b0 + nb * sb;
a0 < b1 && b0 < a1
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_unchecked_with<T: GemmScalar>(
ws: &mut Workspace,
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
b: *const T,
rsb: isize,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
dispatch::execute(
Task {
m,
k,
n,
alpha,
a,
rsa,
csa,
b,
rsb,
csb,
beta,
c,
rsc,
csc,
},
par,
ws,
);
}
}