#[cfg(feature = "epilogue")]
use crate::Bias;
#[cfg(all(feature = "int8", feature = "epilogue"))]
use crate::RequantScale;
#[cfg(feature = "epilogue")]
use crate::kernel::epilogue::BiasDim;
#[cfg(feature = "epilogue")]
#[inline]
pub fn c_byte_range<T>(cp: *const T, dims: &[(usize, isize)]) -> (usize, usize) {
let sz = core::mem::size_of::<T>() as isize;
if dims.iter().any(|&(d, _)| d == 0) {
let b = cp as usize;
return (b, b);
}
let (mut lo, mut hi): (isize, isize) = (0, 0);
for &(d, s) in dims {
if d <= 1 {
continue; }
let e = (d as isize - 1) * s;
if e < 0 {
lo += e;
} else {
hi += e;
}
}
let base = cp as isize;
((base + lo * sz) as usize, (base + (hi + 1) * sz) as usize)
}
#[cfg(feature = "epilogue")]
#[inline]
pub fn bias_overlaps_c<TC, TB>(
cp: *const TC,
c_dims: &[(usize, isize)],
bias: *const TB,
len: usize,
) -> bool {
let (c_lo, c_hi) = c_byte_range(cp, c_dims);
if c_lo == c_hi || len == 0 {
return false;
}
let b_lo = bias as usize;
let b_hi = b_lo + len * core::mem::size_of::<TB>();
c_lo < b_hi && b_lo < c_hi
}
#[cfg(feature = "epilogue")]
pub fn lower_bias<T>(
bias: Option<Bias<'_, T>>,
m: usize,
n: usize,
cp: *const T,
c_dims: &[(usize, isize)],
) -> (*const T, BiasDim, bool) {
match bias {
None => (core::ptr::null(), BiasDim::PerRow, false),
Some(Bias::PerRow(s)) => {
assert_eq!(
s.len(),
m,
"gemmkit: PerRow bias length ({}) != A.rows ({})",
s.len(),
m
);
if bias_overlaps_c(cp, c_dims, s.as_ptr(), s.len()) {
panic!("gemmkit: bias slice overlaps C");
}
(s.as_ptr(), BiasDim::PerRow, true)
}
Some(Bias::PerCol(s)) => {
assert_eq!(
s.len(),
n,
"gemmkit: PerCol bias length ({}) != B.cols ({})",
s.len(),
n
);
if bias_overlaps_c(cp, c_dims, s.as_ptr(), s.len()) {
panic!("gemmkit: bias slice overlaps C");
}
(s.as_ptr(), BiasDim::PerCol, true)
}
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub fn requant_bias<TC>(
m: usize,
cp: *const TC,
c_dims: &[(usize, isize)],
bias: Option<&[i32]>,
) -> (*const i32, bool) {
match bias {
Some(bias) => {
assert_eq!(
bias.len(),
m,
"gemmkit: requantize bias length ({}) != A.rows ({})",
bias.len(),
m
);
if bias_overlaps_c(cp, c_dims, bias.as_ptr(), bias.len()) {
panic!("gemmkit: requantize bias overlaps C");
}
(bias.as_ptr(), true)
}
None => (core::ptr::null(), false),
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub fn requant_scale<TC>(
m: usize,
cp: *const TC,
c_dims: &[(usize, isize)],
scale: RequantScale<'_>,
) -> (f32, *const f32, bool) {
match scale {
RequantScale::PerTensor(s) => {
assert!(
s.is_finite() && s > 0.0,
"gemmkit: requantize scale ({s}) must be finite and > 0"
);
(s, core::ptr::null(), false)
}
RequantScale::PerRow(scales) => {
assert_eq!(
scales.len(),
m,
"gemmkit: requantize scales length ({}) != A.rows ({})",
scales.len(),
m
);
if bias_overlaps_c(cp, c_dims, scales.as_ptr(), scales.len()) {
panic!("gemmkit: requantize scales overlap C");
}
for &s in scales {
assert!(
s.is_finite() && s > 0.0,
"gemmkit: requantize scale ({s}) must be finite and > 0"
);
}
(0.0, scales.as_ptr(), true)
}
}
}