use std::cell::RefCell;
use std::marker::PhantomData;
use std::mem::MaybeUninit;
use std::ops::{Add, Mul, Range};
use rayon::prelude::*;
use rten_base::byte_cast::Pod;
use rten_base::iter::{MaybeParIter, range_chunks};
use rten_base::num::Identities;
use rten_tensor::prelude::*;
use rten_tensor::{
Alloc, AssumeInit, GlobalAlloc, Matrix, MatrixLayout, MatrixMut, MutLayout, NdLayout, NdTensor,
NdTensorView, OverlapPolicy, Storage,
};
mod errors;
mod im2col;
mod kernels;
mod packing;
mod prepack;
mod tiles;
pub use errors::GemmError;
pub use im2col::{ColOffsets, Im2Col, RowOffsets};
pub use kernels::QuantParams;
use kernels::generic::GenericKernel;
use kernels::{Kernel, MatVecOutput};
use packing::PackingBuffer;
pub use prepack::{PackedAMatrix, PackedBMatrix};
use tiles::OutputTiles;
pub type GemmResult<T = ()> = Result<T, GemmError>;
#[derive(Copy, Clone)]
pub enum GemmInputA<'a, T> {
Unpacked(Matrix<'a, T>),
Packed(&'a PackedAMatrix<T>),
}
impl<T> GemmInputA<'_, T> {
pub fn rows(&self) -> usize {
match self {
Self::Unpacked(m) => m.rows(),
Self::Packed(pm) => pm.rows(),
}
}
pub fn cols(&self) -> usize {
match self {
Self::Unpacked(m) => m.cols(),
Self::Packed(pm) => pm.cols(),
}
}
}
pub trait GemmInT: Copy + Default + Send + Sync + Identities + Pod {}
impl GemmInT for i8 {}
impl GemmInT for u8 {}
impl GemmInT for f32 {}
pub trait GemmOutT:
Copy
+ Default
+ PartialEq
+ Send
+ Sync
+ Mul<Self, Output = Self>
+ Add<Self, Output = Self>
+ Identities
+ Pod
{
}
impl GemmOutT for i32 {}
impl GemmOutT for f32 {}
#[derive(Copy, Clone)]
pub enum GemmInputB<'a, T> {
Unpacked(Matrix<'a, T>),
Packed(&'a PackedBMatrix<T>),
Im2Col(&'a Im2Col<'a, T>),
}
impl<T: Copy + Default> GemmInputB<'_, T> {
pub fn rows(&self) -> usize {
match self {
Self::Unpacked(m) => m.rows(),
Self::Packed(pm) => pm.rows(),
Self::Im2Col(im) => im.rows(),
}
}
pub fn cols(&self) -> usize {
match self {
Self::Unpacked(m) => m.cols(),
Self::Packed(pm) => pm.cols(),
Self::Im2Col(im) => im.cols(),
}
}
}
#[derive(Copy, Clone, PartialEq, Debug)]
pub enum BiasVector<'a, T> {
Column(&'a [T]),
Row(&'a [T]),
}
#[derive(Clone, Copy, Debug)]
enum F32KernelType {
#[allow(unused)]
Generic,
#[cfg(target_arch = "x86_64")]
Fma,
#[cfg(target_arch = "x86_64")]
Avx512,
#[cfg(target_arch = "aarch64")]
ArmNeon,
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
Wasm,
}
#[derive(Clone, Copy, Debug)]
enum Int8KernelType {
#[allow(unused)]
Generic,
#[cfg(target_arch = "x86_64")]
Avx2,
#[cfg(target_arch = "x86_64")]
Avx512,
#[cfg(target_arch = "aarch64")]
ArmI8mm,
#[cfg(target_arch = "aarch64")]
ArmDot,
#[cfg(target_arch = "aarch64")]
ArmNeon,
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
Wasm,
}
pub struct GemmExecutor<LhsT: GemmInT = f32, RhsT: GemmInT = LhsT, OutT: GemmOutT = LhsT> {
kernel: Box<dyn Kernel<LhsT, RhsT, OutT>>,
}
impl<LhsT: GemmInT, RhsT: GemmInT, OutT: GemmOutT> GemmExecutor<LhsT, RhsT, OutT> {
pub fn new() -> Self
where
Self: Default,
{
Self::default()
}
#[allow(dead_code)]
pub fn kernel_name(&self) -> &str {
self.kernel.name()
}
#[allow(unused)]
pub fn prepack_a(&self, a: Matrix<LhsT>) -> PackedAMatrix<LhsT> {
self.prepack_a_in(GlobalAlloc::new(), a)
}
pub fn prepack_a_in<A: Alloc>(&self, alloc: A, a: Matrix<LhsT>) -> PackedAMatrix<LhsT> {
prepack::prepack_a(&*self.kernel, alloc, a)
}
pub fn im2col_col_count_step(&self) -> usize {
self.kernel.im2col_col_count_step()
}
pub fn im2col_row_count_step(&self) -> usize {
self.kernel.im2col_row_count_step()
}
#[allow(unused)]
pub fn prepack_b(&self, b: Matrix<RhsT>) -> PackedBMatrix<RhsT> {
self.prepack_b_in(GlobalAlloc::new(), b)
}
pub fn prepack_b_in<A: Alloc>(&self, alloc: A, b: Matrix<RhsT>) -> PackedBMatrix<RhsT> {
prepack::prepack_b(&*self.kernel, alloc, b)
}
pub fn gemm(
&self,
out_data: &mut [OutT],
a: GemmInputA<LhsT>,
b: GemmInputB<RhsT>,
alpha: f32,
beta: OutT,
bias: Option<BiasVector<OutT>>,
a_quant: Option<QuantParams<LhsT>>,
b_quant: Option<QuantParams<RhsT>>,
) -> GemmResult {
gemm_impl(
&*self.kernel,
unsafe { std::mem::transmute::<&mut [OutT], &mut [MaybeUninit<OutT>]>(out_data) },
a,
b,
alpha,
beta,
bias,
a_quant,
b_quant,
)
.map(|_| ())
}
pub fn gemm_uninit<'a>(
&self,
out_data: &'a mut [MaybeUninit<OutT>],
a: GemmInputA<LhsT>,
b: GemmInputB<RhsT>,
alpha: f32,
bias: Option<BiasVector<OutT>>,
a_quant: Option<QuantParams<LhsT>>,
b_quant: Option<QuantParams<RhsT>>,
) -> GemmResult<&'a mut [OutT]> {
gemm_impl(
&*self.kernel,
out_data,
a,
b,
alpha,
OutT::zero(),
bias,
a_quant,
b_quant,
)
}
pub fn batched_gemm_uninit<'a>(
&self,
out_data: &'a mut [MaybeUninit<OutT>],
a: &[GemmInputA<LhsT>],
b: &[GemmInputB<RhsT>],
alpha: f32,
bias: Option<BiasVector<OutT>>,
a_quant: Option<QuantParams<LhsT>>,
b_quant: Option<QuantParams<RhsT>>,
) -> GemmResult<&'a mut [OutT]> {
if a.len() != b.len() {
return Err(GemmError::BatchSizeMismatch);
}
let out_mat_stride = match (a, b) {
([a, ..], [b, ..]) => a.rows() * b.cols(),
_ => 0,
};
if a.len() * out_mat_stride != out_data.len() {
return Err(GemmError::OutputSizeMismatch);
}
match (a, b) {
([], []) => {
Ok(unsafe { out_data.assume_init() })
}
([a], [b]) => {
self.gemm_uninit(out_data, *a, *b, alpha, bias, a_quant, b_quant)
}
(a, b) => {
a.par_iter()
.zip(b)
.zip(out_data.par_chunks_mut(out_mat_stride))
.try_for_each(|((a_mat, b_mat), out_mat)| {
self.gemm_uninit(out_mat, *a_mat, *b_mat, alpha, bias, a_quant, b_quant)
.map(|_| ())
})?;
Ok(unsafe { out_data.assume_init() })
}
}
}
pub fn may_saturate(&self) -> bool {
self.kernel.may_saturate()
}
fn from_kernel<K: Kernel<LhsT, RhsT, OutT> + 'static>() -> Option<Self> {
K::new().map(|kernel| GemmExecutor {
kernel: Box::new(kernel),
})
}
}
macro_rules! try_kernel {
($hint:expr) => {
if let Some(gemm) = Self::with_kernel($hint) {
return gemm;
}
};
}
trait WithKernel: Sized {
type KernelType;
fn with_kernel(kern_type: Self::KernelType) -> Option<Self>;
fn with_generic_kernel() -> Self;
#[allow(unused)]
fn kernel_types() -> Vec<Self::KernelType>;
}
impl WithKernel for GemmExecutor<f32, f32, f32> {
type KernelType = F32KernelType;
#[allow(dead_code)] fn with_kernel(hint: F32KernelType) -> Option<Self> {
match hint {
#[cfg(target_arch = "x86_64")]
F32KernelType::Avx512 => Self::from_kernel::<kernels::x86_64::Avx512Kernel>(),
#[cfg(target_arch = "x86_64")]
F32KernelType::Fma => Self::from_kernel::<kernels::x86_64::FmaKernel>(),
#[cfg(target_arch = "aarch64")]
F32KernelType::ArmNeon => Self::from_kernel::<kernels::aarch64::ArmNeonKernel>(),
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
F32KernelType::Wasm => Self::from_kernel::<kernels::wasm::WasmKernel>(),
F32KernelType::Generic => Some(Self::with_generic_kernel()),
}
}
fn kernel_types() -> Vec<F32KernelType> {
let mut types = Vec::new();
#[cfg(target_arch = "x86_64")]
{
types.push(F32KernelType::Avx512);
types.push(F32KernelType::Fma);
}
#[cfg(target_arch = "aarch64")]
{
types.push(F32KernelType::ArmNeon);
}
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
{
types.push(F32KernelType::Wasm);
}
types.push(F32KernelType::Generic);
types
}
fn with_generic_kernel() -> Self {
Self::from_kernel::<GenericKernel>().unwrap()
}
}
impl Default for GemmExecutor<f32, f32, f32> {
fn default() -> Self {
#[cfg(target_arch = "x86_64")]
try_kernel!(F32KernelType::Avx512);
#[cfg(target_arch = "x86_64")]
try_kernel!(F32KernelType::Fma);
#[cfg(target_arch = "aarch64")]
try_kernel!(F32KernelType::ArmNeon);
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
try_kernel!(F32KernelType::Wasm);
Self::with_generic_kernel()
}
}
impl WithKernel for GemmExecutor<u8, i8, i32> {
type KernelType = Int8KernelType;
fn with_kernel(hint: Int8KernelType) -> Option<Self> {
match hint {
#[cfg(target_arch = "x86_64")]
Int8KernelType::Avx512 => Self::from_kernel::<kernels::x86_64::Avx512Int8Kernel>(),
#[cfg(target_arch = "x86_64")]
Int8KernelType::Avx2 => Self::from_kernel::<kernels::x86_64::Avx2Int8Kernel>(),
#[cfg(target_arch = "aarch64")]
Int8KernelType::ArmNeon => Self::from_kernel::<kernels::aarch64::ArmInt8MlalKernel>(),
#[cfg(target_arch = "aarch64")]
Int8KernelType::ArmDot => Self::from_kernel::<kernels::aarch64::ArmInt8DotKernel>(),
#[cfg(target_arch = "aarch64")]
Int8KernelType::ArmI8mm => Self::from_kernel::<kernels::aarch64::ArmInt8MMKernel>(),
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
Int8KernelType::Wasm => Self::from_kernel::<kernels::wasm::WasmInt8Kernel>(),
Int8KernelType::Generic => Self::from_kernel::<GenericKernel>(),
}
}
fn kernel_types() -> Vec<Int8KernelType> {
let mut types = Vec::new();
#[cfg(target_arch = "x86_64")]
{
types.push(Int8KernelType::Avx512);
types.push(Int8KernelType::Avx2);
}
#[cfg(target_arch = "aarch64")]
{
types.push(Int8KernelType::ArmI8mm);
types.push(Int8KernelType::ArmDot);
types.push(Int8KernelType::ArmNeon);
}
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
{
types.push(Int8KernelType::Wasm);
}
types.push(Int8KernelType::Generic);
types
}
fn with_generic_kernel() -> Self {
Self::from_kernel::<GenericKernel>().unwrap()
}
}
impl Default for GemmExecutor<u8, i8, i32> {
fn default() -> Self {
#[cfg(target_arch = "x86_64")]
{
try_kernel!(Int8KernelType::Avx512);
try_kernel!(Int8KernelType::Avx2);
}
#[cfg(target_arch = "aarch64")]
{
try_kernel!(Int8KernelType::ArmI8mm);
try_kernel!(Int8KernelType::ArmDot);
try_kernel!(Int8KernelType::ArmNeon);
}
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
{
try_kernel!(Int8KernelType::Wasm);
}
Self::with_generic_kernel()
}
}
fn depth_block_size<RhsT>(a_cols: usize) -> usize {
let max = 1024 / size_of::<RhsT>();
max.min(a_cols)
}
fn col_block_size(b_cols: usize, nr: usize) -> usize {
let parallelism = rayon::current_num_threads();
let lower_bound = 128.min(b_cols);
let unrounded = (b_cols / parallelism).max(lower_bound).min(1024);
unrounded.next_multiple_of(nr)
}
fn row_block_size(a_rows: usize, mr: usize) -> usize {
64.min(a_rows).next_multiple_of(mr)
}
fn gemv<'a, LhsT: GemmInT, RhsT: GemmInT, OutT: GemmOutT>(
kernel: &dyn Kernel<LhsT, RhsT, OutT>,
a: NdTensorView<LhsT, 1>,
b: Matrix<RhsT>,
out_data: &'a mut [MaybeUninit<OutT>],
alpha: f32,
beta: OutT,
bias: Option<BiasVector<OutT>>,
a_quant: Option<QuantParams<LhsT>>,
b_quant: Option<QuantParams<RhsT>>,
) -> &'a mut [OutT] {
let a_cols = a.size(0);
let b_cols = b.cols();
let a = a.to_contiguous();
let a_data = a.data().unwrap();
let b_block_size = b_cols.div_ceil(rayon::current_num_threads()).max(128);
let k_block_size = if b.row_stride() == 1 { 512 } else { 8 };
out_data
.par_chunks_mut(b_block_size)
.enumerate()
.for_each(|(col_block_idx, out_chunk)| {
let col_block =
(col_block_idx * b_block_size)..((col_block_idx + 1) * b_block_size).min(b_cols);
let mut effective_beta = beta;
let b_quant = b_quant.map(|bq| QuantParams {
zero_point: &bq.zero_point[col_block.clone()],
});
for (k_block, a_block) in
range_chunks(0..a_cols, k_block_size).zip(a_data.chunks(k_block_size))
{
let b_block = slice_matrix(b, k_block, col_block.clone());
let mat_vec_out = if effective_beta == OutT::zero() {
MatVecOutput::from_uninit_slice(out_chunk)
} else {
MatVecOutput::from_slice(unsafe { out_chunk.assume_init() }, effective_beta)
};
kernel.gemv_kernel(mat_vec_out, a_block, b_block, alpha, a_quant, b_quant);
effective_beta = OutT::one();
}
let out_chunk = unsafe { out_chunk.assume_init() };
match bias {
Some(BiasVector::Column(bias)) => {
let bias = bias[0];
for x in out_chunk {
*x = *x + bias;
}
}
Some(BiasVector::Row(bias)) => {
let bias_block = &bias[col_block.clone()];
for (x, bias) in out_chunk.iter_mut().zip(bias_block) {
*x = *x + *bias;
}
}
None => {}
}
});
unsafe { out_data.assume_init() }
}
fn slice_matrix<T>(mat: Matrix<T>, rows: Range<usize>, cols: Range<usize>) -> Matrix<T> {
let layout = NdLayout::from_shape_and_strides(
[rows.len(), cols.len()],
[mat.row_stride(), mat.col_stride()],
OverlapPolicy::AllowOverlap,
)
.unwrap();
let offset = rows.start * mat.row_stride() + cols.start * mat.col_stride();
Matrix::from_storage_and_layout(
mat.storage().slice(offset..offset + layout.min_data_len()),
layout,
)
}
fn gemm_impl<'a, LhsT: GemmInT, RhsT: GemmInT, OutT: GemmOutT>(
kernel: &dyn Kernel<LhsT, RhsT, OutT>,
out_data: &'a mut [MaybeUninit<OutT>],
a: GemmInputA<LhsT>,
b: GemmInputB<RhsT>,
alpha: f32,
beta: OutT,
bias: Option<BiasVector<OutT>>,
a_quant: Option<QuantParams<LhsT>>,
b_quant: Option<QuantParams<RhsT>>,
) -> GemmResult<&'a mut [OutT]> {
if a.cols() != b.rows() {
return Err(GemmError::KSizeMismatch);
}
let bias_ok = match bias {
Some(BiasVector::Row(bias)) => bias.len() == b.cols(),
Some(BiasVector::Column(bias)) => bias.len() == a.rows(),
None => true,
};
if !bias_ok {
return Err(GemmError::WrongBiasSize);
}
if let Some(a_quant) = a_quant
&& a_quant.zero_point.len() != a.rows()
{
return Err(GemmError::WrongQuantParamSize);
}
if let Some(b_quant) = b_quant
&& b_quant.zero_point.len() != b.cols()
{
return Err(GemmError::WrongQuantParamSize);
}
let mut output_mat =
MatrixMut::<MaybeUninit<OutT>>::try_from_data([a.rows(), b.cols()], out_data)
.map_err(|_| GemmError::OutputSizeMismatch)?;
if a.rows() == 0 || b.cols() == 0 {
let empty = NdTensor::zeros(output_mat.shape());
return Ok(output_mat.init_from(&empty).into_slice_mut().unwrap());
}
if a.cols() == 0 {
let mut output_mat = if beta == OutT::zero() {
output_mat.fill(MaybeUninit::new(OutT::zero()));
unsafe { output_mat.assume_init() }
} else {
let mut output_mat = unsafe { output_mat.assume_init() };
output_mat.apply(|x| *x * beta);
output_mat
};
if let Some(bias) = bias {
let bias_mat = match bias {
BiasVector::Column(bias) => {
NdTensorView::from_data([a.rows(), 1], bias).broadcast([a.rows(), b.cols()])
}
BiasVector::Row(bias) => {
NdTensorView::from_data([1, b.cols()], bias).broadcast([a.rows(), b.cols()])
}
};
for r in 0..a.rows() {
for c in 0..b.cols() {
let out_el = &mut output_mat[[r, c]];
*out_el = *out_el + bias_mat[[r, c]];
}
}
}
return Ok(output_mat.into_slice_mut().unwrap());
}
if let (1, GemmInputA::Unpacked(a), GemmInputB::Unpacked(b)) = (a.rows(), a, b) {
let output = gemv(
kernel,
a.slice(0),
b,
output_mat.into_slice_mut().unwrap(),
alpha,
beta,
bias,
a_quant,
b_quant,
);
return Ok(output);
}
let output_tiles = OutputTiles::new(output_mat.view_mut(), kernel.mr(), kernel.nr());
let nc = col_block_size(b.cols(), kernel.nr());
let mc = row_block_size(a.rows(), kernel.mr());
let kc = depth_block_size::<RhsT>(a.cols());
if let GemmInputA::Packed(packed) = &a {
packed.validate(kernel, kc)?;
}
if let GemmInputB::Packed(packed) = &b {
packed.validate(kernel, kc)?;
}
thread_local!(static PACKED_A: RefCell<PackingBuffer> = const { RefCell::new(PackingBuffer::new()) });
thread_local!(static PACKED_B: RefCell<PackingBuffer> = const { RefCell::new(PackingBuffer::new()) });
let n_col_blocks = b.cols().div_ceil(nc);
let n_row_blocks = a.rows().div_ceil(mc);
let parallel = rayon::current_num_threads() > 1;
let (mr, nr) = (kernel.mr(), kernel.nr());
(0..n_col_blocks)
.maybe_par_iter(parallel)
.for_each(|col_idx| {
let col_start = col_idx * nc;
let col_end = (col_start + nc).min(b.cols());
let col_range = col_start..col_end;
for (depth_block_idx, depth_range) in range_chunks(0..a.cols(), kc).enumerate() {
let mut thread_local_packed_b: Option<PackingBuffer> = None;
let rhs_block = match b {
GemmInputB::Unpacked(_) | GemmInputB::Im2Col(_) => PACKED_B.with(|cell| {
let mut packed_b = cell.take();
let layout =
kernel.packed_b_layout(depth_range.len(), col_end - col_start, b_quant);
let packed_uninit = packed_b.alloc(layout.size(), layout.align());
match b {
GemmInputB::Unpacked(b) => kernel.pack_b_block(
packed_uninit,
b,
depth_range.clone(),
col_start..col_end,
b_quant,
),
GemmInputB::Im2Col(im) => kernel.pack_im2col(
packed_uninit,
im,
depth_range.clone(),
col_start..col_end,
b_quant.map(|q| q.zero_point[0]),
),
GemmInputB::Packed(_) => unreachable!(),
}
unsafe {
packed_b.set_len(layout.size());
}
thread_local_packed_b = Some(packed_b);
RhsBlock {
data: thread_local_packed_b.as_ref().unwrap().as_bytes(),
panel_stride: layout.panel_stride(),
_marker: PhantomData,
}
}),
GemmInputB::Packed(pm) => pm.block(col_range.clone(), depth_block_idx),
};
let effective_beta = if depth_range.start == 0 {
beta
} else {
OutT::one()
};
(0..n_row_blocks)
.maybe_par_iter(parallel)
.for_each(|row_idx| {
let row_start = row_idx * mc;
let row_end = (row_start + mc).min(a.rows());
let row_range = row_start..row_end;
let mut thread_local_packed_a: Option<PackingBuffer> = None;
let lhs_block = match a {
GemmInputA::Unpacked(a) => PACKED_A.with(|cell| {
let layout = kernel.packed_a_layout(
a,
row_end - row_start,
depth_range.len(),
a_quant,
);
if !layout.must_pack {
return LhsBlock::Unpacked(a);
};
let mut packed_a = cell.take();
let packed_uninit = packed_a.alloc(layout.size(), layout.align());
kernel.pack_a_block(
packed_uninit,
a,
row_start..row_end,
depth_range.clone(),
a_quant,
);
unsafe {
packed_a.set_len(layout.size());
}
thread_local_packed_a = Some(packed_a);
LhsBlock::Packed {
data: thread_local_packed_a.as_ref().unwrap().as_bytes(),
panel_stride: layout.panel_stride(),
}
}),
GemmInputA::Packed(pm) => pm.block(row_range.clone(), depth_block_idx),
};
gemm_block(
kernel,
&output_tiles,
col_start / nr..col_end.div_ceil(nr),
row_start / mr..row_end.div_ceil(mr),
depth_range.clone(),
lhs_block,
rhs_block,
alpha,
effective_beta,
bias,
a_quant,
b_quant,
);
if let Some(packed_a) = thread_local_packed_a {
PACKED_A.with(|cell| cell.replace(packed_a));
}
});
if let Some(packed_b) = thread_local_packed_b {
PACKED_B.with(|cell| cell.replace(packed_b));
}
}
});
let output = unsafe { output_mat.assume_init() };
Ok(output.into_slice_mut().unwrap())
}
#[derive(Copy, Clone)]
enum LhsBlock<'a, T> {
Packed {
data: &'a [u8],
panel_stride: usize,
},
Unpacked(Matrix<'a, T>),
}
#[derive(Copy, Clone)]
struct RhsBlock<'a, T> {
data: &'a [u8],
panel_stride: usize,
_marker: PhantomData<T>,
}
fn gemm_block<LhsT: Sync, RhsT: Sync, OutT: GemmOutT>(
kernel: &dyn Kernel<LhsT, RhsT, OutT>,
output: &OutputTiles<MaybeUninit<OutT>>,
col_tiles: Range<usize>,
row_tiles: Range<usize>,
depth_range: Range<usize>,
a: LhsBlock<LhsT>,
b: RhsBlock<RhsT>,
alpha: f32,
beta: OutT,
bias: Option<BiasVector<OutT>>,
a_quant: Option<QuantParams<LhsT>>,
b_quant: Option<QuantParams<RhsT>>,
) {
let (mr, nr) = (kernel.mr(), kernel.nr());
if let LhsBlock::Unpacked(mat) = &a {
assert!(mat.rows().div_ceil(mr) >= row_tiles.end);
assert!(mat.cols() >= depth_range.end);
}
col_tiles
.enumerate()
.for_each(|(block_col_tile, col_tile)| {
let b_panel_offset = block_col_tile * b.panel_stride;
let b_panel = &b.data[b_panel_offset..b_panel_offset + b.panel_stride];
let b_quant_tile = b_quant.map(|bq| {
let col_range = col_tile * nr..(col_tile * nr + nr).min(bq.zero_point.len());
QuantParams {
zero_point: &bq.zero_point[col_range],
}
});
for (block_row_tile, row_tile) in row_tiles.clone().enumerate() {
let out_tile = unsafe { output.tile(row_tile, col_tile) };
let a_quant_tile = a_quant.map(|aq| {
let row_range = row_tile * mr..(row_tile * mr + mr).min(aq.zero_point.len());
QuantParams {
zero_point: &aq.zero_point[row_range],
}
});
let kernel_lhs = match a {
LhsBlock::Packed {
data,
panel_stride,
} => {
let a_panel_offset = block_row_tile * panel_stride;
let a_panel = &data[a_panel_offset..a_panel_offset + panel_stride];
kernels::Lhs::Packed(a_panel)
}
LhsBlock::Unpacked(mat) => {
let storage = mat.storage();
let offset =
row_tile * mr * mat.row_stride() + depth_range.start * mat.col_stride();
kernels::Lhs::Unpacked {
data: unsafe { storage.as_ptr().add(offset) },
len: storage.len().saturating_sub(offset),
row_stride: mat.row_stride(),
}
}
};
unsafe {
kernel.kernel(
out_tile.ptr as *mut OutT,
out_tile.row_stride,
kernel_lhs,
b_panel,
out_tile.used_rows,
out_tile.used_cols,
depth_range.len(),
alpha,
beta,
a_quant_tile,
b_quant_tile,
);
}
if depth_range.start == 0 {
let out_ptr = out_tile.ptr as *mut OutT;
match bias {
Some(BiasVector::Column(bias)) => {
for row in 0..out_tile.used_rows {
for col in 0..out_tile.used_cols {
unsafe {
let out_el = out_ptr.add(row * out_tile.row_stride + col);
*out_el =
*out_el + *bias.get_unchecked(row_tile * mr + row);
}
}
}
}
Some(BiasVector::Row(bias)) => {
for row in 0..out_tile.used_rows {
for col in 0..out_tile.used_cols {
unsafe {
let out_el = out_ptr.add(row * out_tile.row_stride + col);
*out_el =
*out_el + *bias.get_unchecked(col_tile * nr + col);
}
}
}
}
None => {}
}
}
}
});
}
mod reduced_range_rng;
#[doc(hidden)]
pub use reduced_range_rng::ReducedRangeRng;
#[cfg(test)]
mod tests;