use super::*;
#[cfg(feature = "epilogue")]
use crate::kernel::epilogue::BiasDim;
#[cfg(feature = "int8")]
pub fn gemm_i8(
alpha: i32,
a: MatRef<'_, i8>,
b: MatRef<'_, i8>,
beta: i32,
c: MatMut<'_, i32>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| gemm_i8_with(ws, alpha, a, b, beta, c, par));
}
#[cfg(feature = "int8")]
pub fn gemm_i8_with(
ws: &mut Workspace,
alpha: i32,
a: MatRef<'_, i8>,
b: MatRef<'_, i8>,
beta: i32,
c: MatMut<'_, i32>,
par: Parallelism,
) {
validate_gemm_views(&a, &b, &c);
unsafe {
dispatch::execute_int(
dispatch::IntTask {
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,
},
par,
ws,
);
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_i8_unchecked(
m: usize,
k: usize,
n: usize,
alpha: i32,
a: *const i8,
rsa: isize,
csa: isize,
b: *const i8,
rsb: isize,
csb: isize,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
dispatch::execute_int(
dispatch::IntTask {
m,
k,
n,
alpha,
a,
rsa,
csa,
b,
rsb,
csb,
beta,
c,
rsc,
csc,
},
par,
ws,
);
});
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_i8_unchecked_with(
ws: &mut Workspace,
m: usize,
k: usize,
n: usize,
alpha: i32,
a: *const i8,
rsa: isize,
csa: isize,
b: *const i8,
rsb: isize,
csb: isize,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
dispatch::execute_int(
dispatch::IntTask {
m,
k,
n,
alpha,
a,
rsa,
csa,
b,
rsb,
csb,
beta,
c,
rsc,
csc,
},
par,
ws,
);
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
#[derive(Copy, Clone)]
pub enum RequantScale<'a> {
PerTensor(f32),
PerRow(&'a [f32]),
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub struct Requantize<'a> {
pub scale: RequantScale<'a>,
pub zero_point: i32,
pub bias: Option<&'a [i32]>,
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub fn gemm_i8_requant(
a: MatRef<'_, i8>,
b: MatRef<'_, i8>,
req: Requantize<'_>,
c: MatMut<'_, i8>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| gemm_i8_requant_with(ws, a, b, req, c, par));
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
fn requant_bias<TO>(a_rows: usize, c: &MatMut<'_, TO>, bias: Option<&[i32]>) -> (*const i32, bool) {
crate::adapter::requant_bias(a_rows, c.data.as_ptr(), &[(c.data.len(), 1)], bias)
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
fn requant_scale<TO>(
a_rows: usize,
c: &MatMut<'_, TO>,
scale: RequantScale<'_>,
) -> (f32, *const f32, bool) {
crate::adapter::requant_scale(a_rows, c.data.as_ptr(), &[(c.data.len(), 1)], scale)
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub fn gemm_i8_requant_with(
ws: &mut Workspace,
a: MatRef<'_, i8>,
b: MatRef<'_, i8>,
req: Requantize<'_>,
c: MatMut<'_, i8>,
par: Parallelism,
) {
validate_gemm_views(&a, &b, &c);
let (scale, row_scales, has_row_scales) = requant_scale(a.rows, &c, req.scale);
assert!(
(-128..=127).contains(&req.zero_point),
"gemmkit: requantize zero_point ({}) out of i8 range [-128, 127]",
req.zero_point
);
let (bias_ptr, has_bias) = requant_bias(a.rows, &c, req.bias);
unsafe {
dispatch::execute_int_requant(
dispatch::RequantTask {
m: a.rows,
k: a.cols,
n: b.cols,
a: a.data.as_ptr(),
rsa: a.rs,
csa: a.cs,
b: b.data.as_ptr(),
rsb: b.rs,
csb: b.cs,
c: c.data.as_mut_ptr(),
rsc: c.rs,
csc: c.cs,
scale,
row_scales,
has_row_scales,
zp: req.zero_point,
bias: bias_ptr,
has_bias,
bias_dim: BiasDim::PerRow,
},
par,
ws,
);
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_i8_requant_unchecked(
m: usize,
k: usize,
n: usize,
a: *const i8,
rsa: isize,
csa: isize,
b: *const i8,
rsb: isize,
csb: isize,
scale: f32,
row_scales: *const f32,
has_row_scales: bool,
zero_point: i32,
bias: *const i32,
has_bias: bool,
c: *mut i8,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
gemm_i8_requant_unchecked_with(
ws,
m,
k,
n,
a,
rsa,
csa,
b,
rsb,
csb,
scale,
row_scales,
has_row_scales,
zero_point,
bias,
has_bias,
c,
rsc,
csc,
par,
);
});
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_i8_requant_unchecked_with(
ws: &mut Workspace,
m: usize,
k: usize,
n: usize,
a: *const i8,
rsa: isize,
csa: isize,
b: *const i8,
rsb: isize,
csb: isize,
scale: f32,
row_scales: *const f32,
has_row_scales: bool,
zero_point: i32,
bias: *const i32,
has_bias: bool,
c: *mut i8,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
dispatch::execute_int_requant(
dispatch::RequantTask {
m,
k,
n,
a,
rsa,
csa,
b,
rsb,
csb,
c,
rsc,
csc,
scale,
row_scales,
has_row_scales,
zp: zero_point,
bias,
has_bias,
bias_dim: BiasDim::PerRow,
},
par,
ws,
);
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub fn gemm_i8_requant_u8(
a: MatRef<'_, i8>,
b: MatRef<'_, i8>,
req: Requantize<'_>,
c: MatMut<'_, u8>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| gemm_i8_requant_u8_with(ws, a, b, req, c, par));
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
pub fn gemm_i8_requant_u8_with(
ws: &mut Workspace,
a: MatRef<'_, i8>,
b: MatRef<'_, i8>,
req: Requantize<'_>,
c: MatMut<'_, u8>,
par: Parallelism,
) {
validate_gemm_views(&a, &b, &c);
let (scale, row_scales, has_row_scales) = requant_scale(a.rows, &c, req.scale);
assert!(
(0..=255).contains(&req.zero_point),
"gemmkit: requantize zero_point ({}) out of u8 range [0, 255]",
req.zero_point
);
let (bias_ptr, has_bias) = requant_bias(a.rows, &c, req.bias);
unsafe {
dispatch::execute_int_requant(
dispatch::RequantTask {
m: a.rows,
k: a.cols,
n: b.cols,
a: a.data.as_ptr(),
rsa: a.rs,
csa: a.cs,
b: b.data.as_ptr(),
rsb: b.rs,
csb: b.cs,
c: c.data.as_mut_ptr(),
rsc: c.rs,
csc: c.cs,
scale,
row_scales,
has_row_scales,
zp: req.zero_point,
bias: bias_ptr,
has_bias,
bias_dim: BiasDim::PerRow,
},
par,
ws,
);
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_i8_requant_u8_unchecked(
m: usize,
k: usize,
n: usize,
a: *const i8,
rsa: isize,
csa: isize,
b: *const i8,
rsb: isize,
csb: isize,
scale: f32,
row_scales: *const f32,
has_row_scales: bool,
zero_point: i32,
bias: *const i32,
has_bias: bool,
c: *mut u8,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
gemm_i8_requant_u8_unchecked_with(
ws,
m,
k,
n,
a,
rsa,
csa,
b,
rsb,
csb,
scale,
row_scales,
has_row_scales,
zero_point,
bias,
has_bias,
c,
rsc,
csc,
par,
);
});
}
}
#[cfg(all(feature = "int8", feature = "epilogue"))]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_i8_requant_u8_unchecked_with(
ws: &mut Workspace,
m: usize,
k: usize,
n: usize,
a: *const i8,
rsa: isize,
csa: isize,
b: *const i8,
rsb: isize,
csb: isize,
scale: f32,
row_scales: *const f32,
has_row_scales: bool,
zero_point: i32,
bias: *const i32,
has_bias: bool,
c: *mut u8,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
dispatch::execute_int_requant(
dispatch::RequantTask {
m,
k,
n,
a,
rsa,
csa,
b,
rsb,
csb,
c,
rsc,
csc,
scale,
row_scales,
has_row_scales,
zp: zero_point,
bias,
has_bias,
bias_dim: BiasDim::PerRow,
},
par,
ws,
);
}
}