#[cfg(feature = "epilogue")]
use super::fused::{Activation, Bias};
use super::*;
#[cfg(feature = "epilogue")]
use crate::dispatch::FusedScalar;
use crate::dispatch::GemmProblem;
#[cfg(feature = "epilogue")]
use crate::kernel::epilogue::BiasDim;
use alloc::vec::Vec;
#[allow(clippy::too_many_arguments)]
fn check_batched_view<T>(
data: &[T],
rows: usize,
cols: usize,
rs: isize,
cs: isize,
batch: usize,
batch_stride: isize,
name: &str,
) -> usize {
let e = match extent(rows, cols, rs, cs) {
Some(e) => e,
None => panic!(
"gemmkit: {name} view has negative strides or is too large to address; use the unchecked API"
),
};
let last_base = if batch <= 1 {
0
} else {
if batch_stride < 0 {
panic!("gemmkit: {name} batch stride ({batch_stride}) must be non-negative");
}
(batch - 1).saturating_mul(batch_stride as usize)
};
let need = last_base.saturating_add(e);
if need > data.len() {
panic!(
"gemmkit: {name} batched view ({batch}× {rows}x{cols}, batch stride {batch_stride}) \
needs {need} elements but slice has {}",
data.len()
);
}
e
}
#[allow(clippy::too_many_arguments)]
fn validate_batched_views<T>(
batch: usize,
a: &MatRef<'_, T>,
a_batch_stride: isize,
b: &MatRef<'_, T>,
b_batch_stride: isize,
c: &MatMut<'_, T>,
c_batch_stride: isize,
) {
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_batched_view(
a.data,
a.rows,
a.cols,
a.rs,
a.cs,
batch,
a_batch_stride,
"A",
);
check_batched_view(
b.data,
b.rows,
b.cols,
b.rs,
b.cs,
batch,
b_batch_stride,
"B",
);
let c_extent = check_batched_view(
c.data,
c.rows,
c.cols,
c.rs,
c.cs,
batch,
c_batch_stride,
"C",
);
if self_aliases(c.rows, c.cols, c.rs, c.cs) {
panic!(
"gemmkit: batched C element aliases itself (strides {},{} map distinct elements to \
the same memory); C must address each (i,j) uniquely",
c.rs, c.cs
);
}
if batch > 1 && (c_batch_stride as usize) < c_extent {
panic!(
"gemmkit: C batch stride ({c_batch_stride}) must be at least the element extent \
({c_extent}) so the batched C outputs stay disjoint"
);
}
let cp = c.data.as_ptr();
let cl = c.data.len();
if overlaps(cp, cl, a.data.as_ptr(), a.data.len())
|| overlaps(cp, cl, b.data.as_ptr(), b.data.len())
{
panic!("gemmkit: batched C aliases A or B");
}
}
#[allow(clippy::too_many_arguments)]
pub fn gemm_batched<T: GemmScalar>(
batch: usize,
alpha: T,
a: MatRef<'_, T>,
a_batch_stride: isize,
b: MatRef<'_, T>,
b_batch_stride: isize,
beta: T,
c: MatMut<'_, T>,
c_batch_stride: isize,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| {
gemm_batched_with(
ws,
batch,
alpha,
a,
a_batch_stride,
b,
b_batch_stride,
beta,
c,
c_batch_stride,
par,
);
});
}
#[allow(clippy::too_many_arguments)]
pub fn gemm_batched_with<T: GemmScalar>(
ws: &mut Workspace,
batch: usize,
alpha: T,
a: MatRef<'_, T>,
a_batch_stride: isize,
b: MatRef<'_, T>,
b_batch_stride: isize,
beta: T,
c: MatMut<'_, T>,
c_batch_stride: isize,
par: Parallelism,
) {
if batch == 0 {
return;
}
validate_batched_views(
batch,
&a,
a_batch_stride,
&b,
b_batch_stride,
&c,
c_batch_stride,
);
unsafe {
gemm_batched_unchecked_with(
ws,
batch,
a.rows,
a.cols,
b.cols,
alpha,
a.data.as_ptr(),
a.rs,
a.cs,
a_batch_stride,
b.data.as_ptr(),
b.rs,
b.cs,
b_batch_stride,
beta,
c.data.as_mut_ptr(),
c.rs,
c.cs,
c_batch_stride,
par,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub fn gemm_batched_fused<T: FusedScalar>(
batch: usize,
alpha: T,
a: MatRef<'_, T>,
a_batch_stride: isize,
b: MatRef<'_, T>,
b_batch_stride: isize,
beta: T,
c: MatMut<'_, T>,
c_batch_stride: isize,
bias: Option<Bias<'_, T>>,
act: Option<Activation<T>>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| {
gemm_batched_fused_with(
ws,
batch,
alpha,
a,
a_batch_stride,
b,
b_batch_stride,
beta,
c,
c_batch_stride,
bias,
act,
par,
);
});
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub fn gemm_batched_fused_with<T: FusedScalar>(
ws: &mut Workspace,
batch: usize,
alpha: T,
a: MatRef<'_, T>,
a_batch_stride: isize,
b: MatRef<'_, T>,
b_batch_stride: isize,
beta: T,
c: MatMut<'_, T>,
c_batch_stride: isize,
bias: Option<Bias<'_, T>>,
act: Option<Activation<T>>,
par: Parallelism,
) {
if batch == 0 {
return;
}
if bias.is_none() && act.is_none() {
gemm_batched_with(
ws,
batch,
alpha,
a,
a_batch_stride,
b,
b_batch_stride,
beta,
c,
c_batch_stride,
par,
);
return;
}
validate_batched_views(
batch,
&a,
a_batch_stride,
&b,
b_batch_stride,
&c,
c_batch_stride,
);
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 {
crate::special::batched::run_fused(
batch,
a.rows,
a.cols,
b.cols,
alpha,
a.data.as_ptr(),
a.rs,
a.cs,
a_batch_stride,
b.data.as_ptr(),
b.rs,
b.cs,
b_batch_stride,
beta,
c.data.as_mut_ptr(),
c.rs,
c.cs,
c_batch_stride,
epi,
par,
ws,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_batched_fused_unchecked<T: FusedScalar>(
batch: usize,
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
a_batch_stride: isize,
b: *const T,
rsb: isize,
csb: isize,
b_batch_stride: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
c_batch_stride: isize,
bias: *const T,
bias_dim: BiasDim,
has_bias: bool,
act: Option<Activation<T>>,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
gemm_batched_fused_unchecked_with(
ws,
batch,
m,
k,
n,
alpha,
a,
rsa,
csa,
a_batch_stride,
b,
rsb,
csb,
b_batch_stride,
beta,
c,
rsc,
csc,
c_batch_stride,
bias,
bias_dim,
has_bias,
act,
par,
);
});
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_batched_fused_unchecked_with<T: FusedScalar>(
ws: &mut Workspace,
batch: usize,
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
a_batch_stride: isize,
b: *const T,
rsb: isize,
csb: isize,
b_batch_stride: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
c_batch_stride: 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 {
crate::special::batched::run_fused(
batch,
m,
k,
n,
alpha,
a,
rsa,
csa,
a_batch_stride,
b,
rsb,
csb,
b_batch_stride,
beta,
c,
rsc,
csc,
c_batch_stride,
epi,
par,
ws,
);
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_batched_unchecked<T: GemmScalar>(
batch: usize,
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
a_batch_stride: isize,
b: *const T,
rsb: isize,
csb: isize,
b_batch_stride: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
c_batch_stride: isize,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
gemm_batched_unchecked_with(
ws,
batch,
m,
k,
n,
alpha,
a,
rsa,
csa,
a_batch_stride,
b,
rsb,
csb,
b_batch_stride,
beta,
c,
rsc,
csc,
c_batch_stride,
par,
);
});
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_batched_unchecked_with<T: GemmScalar>(
ws: &mut Workspace,
batch: usize,
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
a_batch_stride: isize,
b: *const T,
rsb: isize,
csb: isize,
b_batch_stride: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
c_batch_stride: isize,
par: Parallelism,
) {
unsafe {
crate::special::batched::run(
batch,
m,
k,
n,
alpha,
a,
rsa,
csa,
a_batch_stride,
b,
rsb,
csb,
b_batch_stride,
beta,
c,
rsc,
csc,
c_batch_stride,
par,
ws,
);
}
}
pub unsafe fn gemm_batched_ptr_unchecked<T: GemmScalar>(
problems: &[GemmProblem<T>],
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| crate::special::batched::run_ptr(problems, par, ws));
}
}
pub struct BatchProblem<'a, T> {
pub alpha: T,
pub a: MatRef<'a, T>,
pub b: MatRef<'a, T>,
pub beta: T,
pub c: MatMut<'a, T>,
}
pub fn gemm_batched_slice<T: GemmScalar>(problems: &mut [BatchProblem<'_, T>], par: Parallelism) {
let raw: Vec<GemmProblem<T>> = problems
.iter_mut()
.enumerate()
.map(|(i, p)| {
assert_eq!(
p.a.cols, p.b.rows,
"gemmkit: batch element {i} A.cols ({}) != B.rows ({})",
p.a.cols, p.b.rows
);
assert_eq!(
p.a.rows, p.c.rows,
"gemmkit: batch element {i} A.rows ({}) != C.rows ({})",
p.a.rows, p.c.rows
);
assert_eq!(
p.b.cols, p.c.cols,
"gemmkit: batch element {i} B.cols ({}) != C.cols ({})",
p.b.cols, p.c.cols
);
check_view(p.a.data, p.a.rows, p.a.cols, p.a.rs, p.a.cs, "A");
check_view(p.b.data, p.b.rows, p.b.cols, p.b.rs, p.b.cs, "B");
check_view(p.c.data, p.c.rows, p.c.cols, p.c.rs, p.c.cs, "C");
if self_aliases(p.c.rows, p.c.cols, p.c.rs, p.c.cs) {
panic!(
"gemmkit: batch element {i} C view aliases itself (strides {},{}); C must \
address each (i,j) uniquely",
p.c.rs, p.c.cs
);
}
GemmProblem {
m: p.a.rows,
k: p.a.cols,
n: p.b.cols,
alpha: p.alpha,
a: p.a.data.as_ptr(),
rsa: p.a.rs,
csa: p.a.cs,
b: p.b.data.as_ptr(),
rsb: p.b.rs,
csb: p.b.cs,
beta: p.beta,
c: p.c.data.as_mut_ptr(),
rsc: p.c.rs,
csc: p.c.cs,
}
})
.collect();
workspace::with_thread_pool(|ws| unsafe { crate::special::batched::run_ptr(&raw, par, ws) });
}