use super::*;
use crate::dispatch::FusedScalar;
use crate::kernel::epilogue::{BiasDim, FusedEpi};
#[derive(Copy, Clone, Debug)]
pub enum Bias<'a, T> {
PerRow(&'a [T]),
PerCol(&'a [T]),
}
#[derive(Copy, Clone, Debug)]
pub enum Activation<T> {
Relu,
LeakyRelu(T),
}
#[allow(clippy::too_many_arguments)]
pub fn gemm_fused<T: FusedScalar>(
alpha: T,
a: MatRef<'_, T>,
b: MatRef<'_, T>,
beta: T,
c: MatMut<'_, T>,
bias: Option<Bias<'_, T>>,
act: Option<Activation<T>>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| gemm_fused_with(ws, alpha, a, b, beta, c, bias, act, par));
}
#[allow(clippy::too_many_arguments)]
pub fn gemm_fused_with<T: FusedScalar>(
ws: &mut Workspace,
alpha: T,
a: MatRef<'_, T>,
b: MatRef<'_, T>,
beta: T,
c: MatMut<'_, T>,
bias: Option<Bias<'_, T>>,
act: Option<Activation<T>>,
par: Parallelism,
) {
if bias.is_none() && act.is_none() {
gemm_with(ws, alpha, a, b, beta, c, par);
return;
}
validate_gemm_views(&a, &b, &c);
validate_bias(&bias, a.rows, b.cols, &c);
if let Some(Activation::LeakyRelu(s)) = &act {
assert!(T::finite(*s), "gemmkit: LeakyRelu slope must be finite");
}
let epi = to_fused_epi(bias, act);
unsafe {
dispatch::execute_fused(
Task {
m: a.rows,
k: a.cols,
n: b.cols,
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,
},
epi,
par,
ws,
);
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_fused_unchecked<T: FusedScalar>(
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,
bias: *const T,
bias_dim: BiasDim,
has_bias: bool,
act: Option<Activation<T>>,
par: Parallelism,
) {
let epi = to_fused_epi_raw(bias, bias_dim, has_bias, act);
unsafe {
fused_unchecked_impl(
None, m, k, n, alpha, a, rsa, csa, b, rsb, csb, beta, c, rsc, csc, epi, par,
);
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_fused_unchecked_with<T: FusedScalar>(
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,
bias: *const T,
bias_dim: BiasDim,
has_bias: bool,
act: Option<Activation<T>>,
par: Parallelism,
) {
let epi = to_fused_epi_raw(bias, bias_dim, has_bias, act);
unsafe {
fused_unchecked_impl(
Some(ws),
m,
k,
n,
alpha,
a,
rsa,
csa,
b,
rsb,
csb,
beta,
c,
rsc,
csc,
epi,
par,
);
}
}
#[allow(clippy::too_many_arguments)]
unsafe fn fused_unchecked_impl<T: FusedScalar>(
ws: Option<&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,
epi: FusedEpi<T>,
par: Parallelism,
) {
let task = Task {
m,
k,
n,
alpha,
a,
rsa,
csa,
b,
rsb,
csb,
beta,
c,
rsc,
csc,
};
unsafe {
match ws {
Some(ws) => dispatch::execute_fused(task, epi, par, ws),
None => workspace::with_thread_pool(|ws| dispatch::execute_fused(task, epi, par, ws)),
}
}
}