use super::*;
#[cfg(feature = "epilogue")]
use crate::dispatch::FusedScalar;
use crate::dispatch::PackedConsume;
#[cfg(feature = "epilogue")]
use crate::kernel::epilogue::{BiasDim, BiasSpec, FusedEpi};
#[cfg(feature = "int8")]
use alloc::vec;
use alloc::vec::Vec;
pub struct PackedRhs<T> {
buf: Vec<T>,
k: usize,
n: usize,
nr: usize,
kc: usize,
nc: usize,
}
impl<T> PackedRhs<T> {
pub fn rows(&self) -> usize {
self.k
}
pub fn cols(&self) -> usize {
self.n
}
}
pub fn prepack_rhs<T: GemmScalar>(b: MatRef<'_, T>) -> PackedRhs<T> {
check_view(b.data, b.rows, b.cols, b.rs, b.cs, "B");
unsafe { prepack_rhs_unchecked(b.data.as_ptr(), b.rs, b.cs, b.rows, b.cols) }
}
pub unsafe fn prepack_rhs_unchecked<T: GemmScalar>(
b: *const T,
rsb: isize,
csb: isize,
k: usize,
n: usize,
) -> PackedRhs<T> {
let (mr, nr) = <T as GemmScalar>::rhs_tile();
if k == 0 || n == 0 {
return PackedRhs {
buf: Vec::new(),
k,
n,
nr,
kc: 1,
nc: nr,
};
}
let lhs_size = core::mem::size_of::<T>().max(1);
let dodge_tiny = crate::tuning::tiny_block_dim().saturating_add(1);
let blk = crate::cache::topology().blocking(mr, nr, lhs_size, dodge_tiny, n, k);
let kc = if T::OUT_IS_ACC {
blk.kc.max(1)
} else {
k.max(1)
};
let nc = blk.nc.next_multiple_of(nr).max(nr);
let k_pad = k.next_multiple_of(<T as GemmScalar>::rhs_depth_multiple());
let total = n
.div_ceil(nr)
.checked_mul(nr)
.and_then(|v| v.checked_mul(k_pad))
.unwrap_or_else(|| {
panic!("gemmkit: prepacked RHS of {k}x{n} is too large; the pack buffer size overflows usize")
});
let mut buf: Vec<T> = Vec::with_capacity(total);
if total > 0 {
unsafe {
buf.set_len(total);
T::pack_rhs_full(buf.as_mut_ptr(), b, rsb, csb, k, n, kc, nc, nr);
}
}
PackedRhs {
buf,
k,
n,
nr,
kc,
nc,
}
}
pub fn gemm_packed_b<T: GemmScalar>(
alpha: T,
a: MatRef<'_, T>,
packed: &PackedRhs<T>,
beta: T,
c: MatMut<'_, T>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| gemm_packed_b_with(ws, alpha, a, packed, beta, c, par));
}
pub fn gemm_packed_b_with<T: GemmScalar>(
ws: &mut Workspace,
alpha: T,
a: MatRef<'_, T>,
packed: &PackedRhs<T>,
beta: T,
c: MatMut<'_, T>,
par: Parallelism,
) {
assert_eq!(
a.cols, packed.k,
"gemmkit: A.cols ({}) != packed B.rows ({})",
a.cols, packed.k
);
assert_eq!(
packed.n, c.cols,
"gemmkit: packed B.cols ({}) != C.cols ({})",
packed.n, c.cols
);
assert_eq!(
a.rows, c.rows,
"gemmkit: A.rows ({}) != C.rows ({})",
a.rows, c.rows
);
check_view(a.data, a.rows, a.cols, a.rs, a.cs, "A");
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();
let cl = c.data.len();
if overlaps(cp, cl, a.data.as_ptr(), a.data.len()) {
panic!("gemmkit: C aliases A");
}
unsafe {
gemm_packed_b_unchecked_with(
ws,
alpha,
a.rows,
a.data.as_ptr(),
a.rs,
a.cs,
packed,
beta,
c.data.as_mut_ptr(),
c.rs,
c.cs,
par,
);
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_packed_b_unchecked<T: GemmScalar>(
alpha: T,
m: usize,
a: *const T,
rsa: isize,
csa: isize,
packed: &PackedRhs<T>,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
gemm_packed_b_unchecked_with(ws, alpha, m, a, rsa, csa, packed, beta, c, rsc, csc, par);
});
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_packed_b_unchecked_with<T: GemmScalar>(
ws: &mut Workspace,
alpha: T,
m: usize,
a: *const T,
rsa: isize,
csa: isize,
packed: &PackedRhs<T>,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
par: Parallelism,
) {
assert!(
csc.unsigned_abs() >= rsc.unsigned_abs(),
"gemmkit: gemm_packed_b requires column-major-ish C (|csc| >= |rsc|); a row-major C \
would swap A/B and invalidate the prepacked RHS — use gemm() for that layout"
);
unsafe {
dispatch::execute_packed(
PackedConsume {
m,
k: packed.k,
n: packed.n,
alpha,
a,
rsa,
csa,
packed: packed.buf.as_ptr(),
nr: packed.nr,
kc: packed.kc,
nc: packed.nc,
beta,
c,
rsc,
csc,
},
par,
ws,
);
}
}
#[cfg(feature = "int8")]
pub fn prepack_rhs_i8(b: MatRef<'_, i8>) -> PackedRhs<i8> {
check_view(b.data, b.rows, b.cols, b.rs, b.cs, "B");
unsafe { prepack_rhs_i8_unchecked(b.data.as_ptr(), b.rs, b.cs, b.rows, b.cols) }
}
#[cfg(feature = "int8")]
pub unsafe fn prepack_rhs_i8_unchecked(
b: *const i8,
rsb: isize,
csb: isize,
k: usize,
n: usize,
) -> PackedRhs<i8> {
let (mr, nr) = dispatch::i8_rhs_tile();
if k == 0 || n == 0 {
return PackedRhs {
buf: Vec::new(),
k,
n,
nr,
kc: 1,
nc: nr,
};
}
let dodge_tiny = crate::tuning::tiny_block_dim().saturating_add(1);
let blk = crate::cache::topology().blocking(mr, nr, 1, dodge_tiny, n, k);
let depth_multiple = dispatch::i8_rhs_depth_multiple();
let kc = if depth_multiple > 1 {
k.max(1)
} else {
blk.kc.max(1)
};
let nc = blk.nc.next_multiple_of(nr).max(nr);
let k_pad = k.next_multiple_of(depth_multiple);
let total = n
.div_ceil(nr)
.checked_mul(nr)
.and_then(|v| v.checked_mul(k_pad))
.unwrap_or_else(|| {
panic!("gemmkit: prepacked RHS of {k}x{n} is too large; the pack buffer size overflows usize")
});
let mut buf = vec![0i8; total];
if total > 0 {
unsafe {
dispatch::pack_rhs_full_i8(buf.as_mut_ptr(), b, rsb, csb, k, n, kc, nc, nr);
}
}
PackedRhs {
buf,
k,
n,
nr,
kc,
nc,
}
}
#[cfg(feature = "int8")]
pub fn gemm_i8_packed_b(
alpha: i32,
a: MatRef<'_, i8>,
packed: &PackedRhs<i8>,
beta: i32,
c: MatMut<'_, i32>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| gemm_i8_packed_b_with(ws, alpha, a, packed, beta, c, par));
}
#[cfg(feature = "int8")]
pub fn gemm_i8_packed_b_with(
ws: &mut Workspace,
alpha: i32,
a: MatRef<'_, i8>,
packed: &PackedRhs<i8>,
beta: i32,
c: MatMut<'_, i32>,
par: Parallelism,
) {
assert_eq!(
a.cols, packed.k,
"gemmkit: A.cols ({}) != packed B.rows ({})",
a.cols, packed.k
);
assert_eq!(
packed.n, c.cols,
"gemmkit: packed B.cols ({}) != C.cols ({})",
packed.n, c.cols
);
assert_eq!(
a.rows, c.rows,
"gemmkit: A.rows ({}) != C.rows ({})",
a.rows, c.rows
);
check_view(a.data, a.rows, a.cols, a.rs, a.cs, "A");
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
);
}
if overlaps_bytes(
c.data.as_ptr() as *const u8,
c.data.len(),
core::mem::size_of::<i32>(),
a.data.as_ptr() as *const u8,
a.data.len(),
core::mem::size_of::<i8>(),
) {
panic!("gemmkit: C aliases A");
}
unsafe {
gemm_i8_packed_b_unchecked_with(
ws,
alpha,
a.rows,
a.data.as_ptr(),
a.rs,
a.cs,
packed,
beta,
c.data.as_mut_ptr(),
c.rs,
c.cs,
par,
);
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_i8_packed_b_unchecked(
alpha: i32,
m: usize,
a: *const i8,
rsa: isize,
csa: isize,
packed: &PackedRhs<i8>,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
gemm_i8_packed_b_unchecked_with(
ws, alpha, m, a, rsa, csa, packed, beta, c, rsc, csc, par,
);
});
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_i8_packed_b_unchecked_with(
ws: &mut Workspace,
alpha: i32,
m: usize,
a: *const i8,
rsa: isize,
csa: isize,
packed: &PackedRhs<i8>,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
par: Parallelism,
) {
assert!(
csc.unsigned_abs() >= rsc.unsigned_abs(),
"gemmkit: gemm_packed_b requires column-major-ish C (|csc| >= |rsc|); a row-major C \
would swap A/B and invalidate the prepacked RHS — use gemm() for that layout"
);
unsafe {
dispatch::execute_int_packed(
dispatch::IntPackedConsume {
m,
k: packed.k,
n: packed.n,
alpha,
a,
rsa,
csa,
packed: packed.buf.as_ptr(),
nr: packed.nr,
kc: packed.kc,
nc: packed.nc,
beta,
c,
rsc,
csc,
},
par,
ws,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub fn gemm_packed_b_fused<T: FusedScalar>(
alpha: T,
a: MatRef<'_, T>,
packed: &PackedRhs<T>,
beta: T,
c: MatMut<'_, T>,
bias: Option<Bias<'_, T>>,
act: Option<Activation<T>>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| {
gemm_packed_b_fused_with(ws, alpha, a, packed, beta, c, bias, act, par)
});
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub fn gemm_packed_b_fused_with<T: FusedScalar>(
ws: &mut Workspace,
alpha: T,
a: MatRef<'_, T>,
packed: &PackedRhs<T>,
beta: T,
c: MatMut<'_, T>,
bias: Option<Bias<'_, T>>,
act: Option<Activation<T>>,
par: Parallelism,
) {
assert_eq!(
a.cols, packed.k,
"gemmkit: A.cols ({}) != packed B.rows ({})",
a.cols, packed.k
);
assert_eq!(
packed.n, c.cols,
"gemmkit: packed B.cols ({}) != C.cols ({})",
packed.n, c.cols
);
assert_eq!(
a.rows, c.rows,
"gemmkit: A.rows ({}) != C.rows ({})",
a.rows, c.rows
);
check_view(a.data, a.rows, a.cols, a.rs, a.cs, "A");
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();
let cl = c.data.len();
if overlaps(cp, cl, a.data.as_ptr(), a.data.len()) {
panic!("gemmkit: C aliases A");
}
validate_bias(&bias, a.rows, packed.n, &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 {
packed_b_fused_impl(
Some(ws),
alpha,
a.rows,
a.data.as_ptr(),
a.rs,
a.cs,
packed,
beta,
c.data.as_mut_ptr(),
c.rs,
c.cs,
epi,
par,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_packed_b_fused_unchecked<T: FusedScalar>(
alpha: T,
m: usize,
a: *const T,
rsa: isize,
csa: isize,
packed: &PackedRhs<T>,
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 {
packed_b_fused_impl(
None, alpha, m, a, rsa, csa, packed, beta, c, rsc, csc, epi, par,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_packed_b_fused_unchecked_with<T: FusedScalar>(
ws: &mut Workspace,
alpha: T,
m: usize,
a: *const T,
rsa: isize,
csa: isize,
packed: &PackedRhs<T>,
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 {
packed_b_fused_impl(
Some(ws),
alpha,
m,
a,
rsa,
csa,
packed,
beta,
c,
rsc,
csc,
epi,
par,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
unsafe fn packed_b_fused_impl<T: FusedScalar>(
ws: Option<&mut Workspace>,
alpha: T,
m: usize,
a: *const T,
rsa: isize,
csa: isize,
packed: &PackedRhs<T>,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
epi: FusedEpi<T>,
par: Parallelism,
) {
assert!(
csc.unsigned_abs() >= rsc.unsigned_abs(),
"gemmkit: gemm_packed_b requires column-major-ish C (|csc| >= |rsc|); a row-major C \
would swap A/B and invalidate the prepacked RHS — use gemm() for that layout"
);
let req = PackedConsume {
m,
k: packed.k,
n: packed.n,
alpha,
a,
rsa,
csa,
packed: packed.buf.as_ptr(),
nr: packed.nr,
kc: packed.kc,
nc: packed.nc,
beta,
c,
rsc,
csc,
};
unsafe {
match ws {
Some(ws) => dispatch::execute_packed_fused(req, epi, par, ws),
None => {
workspace::with_thread_pool(|ws| dispatch::execute_packed_fused(req, epi, par, ws))
}
}
}
}
pub struct PackedLhs<T> {
buf: Vec<T>,
m: usize,
k: usize,
nr: usize,
kc: usize,
nc: usize,
}
impl<T> PackedLhs<T> {
pub fn rows(&self) -> usize {
self.m
}
pub fn cols(&self) -> usize {
self.k
}
}
pub fn prepack_lhs<T: GemmScalar>(a: MatRef<'_, T>) -> PackedLhs<T> {
check_view(a.data, a.rows, a.cols, a.rs, a.cs, "A");
unsafe { prepack_lhs_unchecked(a.data.as_ptr(), a.rs, a.cs, a.rows, a.cols) }
}
pub unsafe fn prepack_lhs_unchecked<T: GemmScalar>(
a: *const T,
rsa: isize,
csa: isize,
m: usize,
k: usize,
) -> PackedLhs<T> {
let packed = unsafe { prepack_rhs_unchecked(a, csa, rsa, k, m) };
PackedLhs {
buf: packed.buf,
m: packed.n,
k: packed.k,
nr: packed.nr,
kc: packed.kc,
nc: packed.nc,
}
}
pub fn gemm_packed_a<T: GemmScalar>(
alpha: T,
packed: &PackedLhs<T>,
b: MatRef<'_, T>,
beta: T,
c: MatMut<'_, T>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| gemm_packed_a_with(ws, alpha, packed, b, beta, c, par));
}
pub fn gemm_packed_a_with<T: GemmScalar>(
ws: &mut Workspace,
alpha: T,
packed: &PackedLhs<T>,
b: MatRef<'_, T>,
beta: T,
c: MatMut<'_, T>,
par: Parallelism,
) {
assert_eq!(
packed.k, b.rows,
"gemmkit: packed A.cols ({}) != B.rows ({})",
packed.k, b.rows
);
assert_eq!(
packed.m, c.rows,
"gemmkit: packed A.rows ({}) != C.rows ({})",
packed.m, c.rows
);
assert_eq!(
b.cols, c.cols,
"gemmkit: B.cols ({}) != C.cols ({})",
b.cols, c.cols
);
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();
let cl = c.data.len();
if overlaps(cp, cl, b.data.as_ptr(), b.data.len()) {
panic!("gemmkit: C aliases B");
}
unsafe {
gemm_packed_a_unchecked_with(
ws,
alpha,
packed,
b.cols,
b.data.as_ptr(),
b.rs,
b.cs,
beta,
c.data.as_mut_ptr(),
c.rs,
c.cs,
par,
);
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_packed_a_unchecked<T: GemmScalar>(
alpha: T,
packed: &PackedLhs<T>,
n: usize,
b: *const T,
rsb: isize,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
gemm_packed_a_unchecked_with(ws, alpha, packed, n, b, rsb, csb, beta, c, rsc, csc, par);
});
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_packed_a_unchecked_with<T: GemmScalar>(
ws: &mut Workspace,
alpha: T,
packed: &PackedLhs<T>,
n: usize,
b: *const T,
rsb: isize,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
par: Parallelism,
) {
assert!(
csc.unsigned_abs() <= rsc.unsigned_abs(),
"gemmkit: gemm_packed_a requires row-major-ish C (|csc| <= |rsc|); a column-major C \
would keep A in the LHS role and invalidate the prepacked LHS — use gemm() for that layout"
);
unsafe {
dispatch::execute_packed(
PackedConsume {
m: n,
k: packed.k,
n: packed.m,
alpha,
a: b,
rsa: csb,
csa: rsb,
packed: packed.buf.as_ptr(),
nr: packed.nr,
kc: packed.kc,
nc: packed.nc,
beta,
c,
rsc: csc,
csc: rsc,
},
par,
ws,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub fn gemm_packed_a_fused<T: FusedScalar>(
alpha: T,
packed: &PackedLhs<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_packed_a_fused_with(ws, alpha, packed, b, beta, c, bias, act, par)
});
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub fn gemm_packed_a_fused_with<T: FusedScalar>(
ws: &mut Workspace,
alpha: T,
packed: &PackedLhs<T>,
b: MatRef<'_, T>,
beta: T,
c: MatMut<'_, T>,
bias: Option<Bias<'_, T>>,
act: Option<Activation<T>>,
par: Parallelism,
) {
assert_eq!(
packed.k, b.rows,
"gemmkit: packed A.cols ({}) != B.rows ({})",
packed.k, b.rows
);
assert_eq!(
packed.m, c.rows,
"gemmkit: packed A.rows ({}) != C.rows ({})",
packed.m, c.rows
);
assert_eq!(
b.cols, c.cols,
"gemmkit: B.cols ({}) != C.cols ({})",
b.cols, c.cols
);
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();
let cl = c.data.len();
if overlaps(cp, cl, b.data.as_ptr(), b.data.len()) {
panic!("gemmkit: C aliases B");
}
validate_bias(&bias, packed.m, 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 {
packed_a_fused_impl(
Some(ws),
alpha,
packed,
b.cols,
b.data.as_ptr(),
b.rs,
b.cs,
beta,
c.data.as_mut_ptr(),
c.rs,
c.cs,
epi,
par,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_packed_a_fused_unchecked<T: FusedScalar>(
alpha: T,
packed: &PackedLhs<T>,
n: usize,
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 {
packed_a_fused_impl(
None, alpha, packed, n, b, rsb, csb, beta, c, rsc, csc, epi, par,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_packed_a_fused_unchecked_with<T: FusedScalar>(
ws: &mut Workspace,
alpha: T,
packed: &PackedLhs<T>,
n: usize,
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 {
packed_a_fused_impl(
Some(ws),
alpha,
packed,
n,
b,
rsb,
csb,
beta,
c,
rsc,
csc,
epi,
par,
);
}
}
#[cfg(feature = "epilogue")]
#[allow(clippy::too_many_arguments)]
unsafe fn packed_a_fused_impl<T: FusedScalar>(
ws: Option<&mut Workspace>,
alpha: T,
packed: &PackedLhs<T>,
n: usize,
b: *const T,
rsb: isize,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
mut epi: FusedEpi<T>,
par: Parallelism,
) {
assert!(
csc.unsigned_abs() <= rsc.unsigned_abs(),
"gemmkit: gemm_packed_a requires row-major-ish C (|csc| <= |rsc|); a column-major C \
would keep A in the LHS role and invalidate the prepacked LHS — use gemm() for that layout"
);
epi.bias = match epi.bias {
BiasSpec::None => BiasSpec::None,
BiasSpec::Row(p) => BiasSpec::Col(p),
BiasSpec::Col(p) => BiasSpec::Row(p),
};
let req = PackedConsume {
m: n,
k: packed.k,
n: packed.m,
alpha,
a: b,
rsa: csb,
csa: rsb,
packed: packed.buf.as_ptr(),
nr: packed.nr,
kc: packed.kc,
nc: packed.nc,
beta,
c,
rsc: csc,
csc: rsc,
};
unsafe {
match ws {
Some(ws) => dispatch::execute_packed_fused(req, epi, par, ws),
None => {
workspace::with_thread_pool(|ws| dispatch::execute_packed_fused(req, epi, par, ws))
}
}
}
}