use std::marker::PhantomData;
use std::ops::Range;
use rten_base::iter::range_chunks;
use rten_tensor::storage::Alloc;
use rten_tensor::{Matrix, MatrixLayout};
use super::packing::PackingBuffer;
use super::{GemmError, Kernel, LhsBlock, RhsBlock, depth_block_size};
#[derive(Clone)]
struct PackedMatrixBase {
data: PackingBuffer,
panel_size: usize,
depth_block: usize,
depth_block_stride: usize,
panel_stride: usize,
tail_panel_stride: usize,
nm_size: usize,
depth_size: usize,
kernel_name: &'static str,
}
impl PackedMatrixBase {
fn block(&self, nm_range: Range<usize>, depth_block_idx: usize) -> (&[u8], usize) {
assert_eq!(nm_range.start % self.panel_size, 0);
let n_blocks = self.depth_size.div_ceil(self.depth_block);
let panel_stride = if depth_block_idx == n_blocks - 1 {
self.tail_panel_stride
} else {
self.panel_stride
};
let depth_block_offset = depth_block_idx * self.depth_block_stride;
let panel_range = nm_range.start / self.panel_size..nm_range.end.div_ceil(self.panel_size);
let start = depth_block_offset + panel_range.start * panel_stride;
let end = depth_block_offset + panel_range.end * panel_stride;
let data = &self.data.as_bytes()[start..end];
(data, panel_stride)
}
}
#[derive(Clone)]
pub struct PackedAMatrix<T> {
base: PackedMatrixBase,
_marker: PhantomData<T>,
}
impl<T> PackedAMatrix<T> {
pub(super) fn block(&self, rows: Range<usize>, depth_block_idx: usize) -> LhsBlock<'_, T> {
let (data, panel_stride) = self.base.block(rows, depth_block_idx);
LhsBlock::Packed { data, panel_stride }
}
pub fn cols(&self) -> usize {
self.base.depth_size
}
pub fn rows(&self) -> usize {
self.base.nm_size
}
pub fn validate<LhsT, RhsT, OutT>(
&self,
kernel: &dyn Kernel<LhsT, RhsT, OutT>,
depth_block: usize,
) -> Result<(), GemmError> {
if self.base.panel_size != kernel.mr() || self.base.kernel_name != kernel.name() {
return Err(GemmError::PackedDataKernelMismatch);
}
if self.base.depth_block != depth_block {
return Err(GemmError::PackedDataBlockingMismatch);
}
Ok(())
}
pub fn into_vec(self) -> Vec<u32> {
self.base.data.into_vec()
}
}
#[derive(Clone)]
pub struct PackedBMatrix<T> {
base: PackedMatrixBase,
_marker: PhantomData<T>,
}
impl<T> PackedBMatrix<T> {
pub(super) fn block(&self, cols: Range<usize>, depth_block_idx: usize) -> RhsBlock<'_, T> {
let (data, panel_stride) = self.base.block(cols, depth_block_idx);
RhsBlock {
data,
panel_stride,
_marker: PhantomData,
}
}
pub fn cols(&self) -> usize {
self.base.nm_size
}
pub fn rows(&self) -> usize {
self.base.depth_size
}
pub fn validate<LhsT, RhsT, OutT>(
&self,
kernel: &dyn Kernel<LhsT, RhsT, OutT>,
depth_block: usize,
) -> Result<(), GemmError> {
if self.base.panel_size != kernel.nr() || self.base.kernel_name != kernel.name() {
return Err(GemmError::PackedDataKernelMismatch);
}
if self.base.depth_block != depth_block {
return Err(GemmError::PackedDataBlockingMismatch);
}
Ok(())
}
pub fn into_vec(self) -> Vec<u32> {
self.base.data.into_vec()
}
}
pub fn prepack_a<A: Alloc, LhsT, RhsT, OutT>(
kernel: &dyn Kernel<LhsT, RhsT, OutT>,
alloc: A,
a: Matrix<LhsT>,
) -> PackedAMatrix<LhsT> {
let depth_block = depth_block_size::<RhsT>(a.cols(), None);
let layout = kernel.packed_a_layout(a, a.rows(), depth_block, None);
let tail_layout = if !a.cols().is_multiple_of(depth_block) {
Some(kernel.packed_a_layout(a, a.rows(), a.cols() % depth_block, None))
} else {
None
};
assert_eq!(layout.size() % layout.align(), 0);
let n_blocks = a.cols() / depth_block;
let total_size =
(n_blocks * layout.size()) + tail_layout.as_ref().map(|l| l.size()).unwrap_or(0);
let mut data = PackingBuffer::new();
let uninit_data = data.alloc_in(alloc, total_size, layout.align());
for (col_block, block_data) in
range_chunks(0..a.cols(), depth_block).zip(uninit_data.chunks_mut(layout.size()))
{
kernel.pack_a_block(block_data, a, 0..a.rows(), col_block, None);
}
unsafe {
data.set_len(total_size);
}
PackedAMatrix {
base: PackedMatrixBase {
data,
nm_size: a.rows(),
depth_size: a.cols(),
panel_size: kernel.mr(),
depth_block,
panel_stride: layout.panel_stride(),
tail_panel_stride: tail_layout
.map(|tl| tl.panel_stride())
.unwrap_or(layout.panel_stride()),
depth_block_stride: layout.size(),
kernel_name: kernel.name(),
},
_marker: PhantomData,
}
}
pub fn prepack_b<A: Alloc, LhsT, RhsT, OutT>(
kernel: &dyn Kernel<LhsT, RhsT, OutT>,
alloc: A,
b: Matrix<RhsT>,
) -> PackedBMatrix<RhsT> {
let depth_block = depth_block_size::<RhsT>(b.rows(), None);
let layout = kernel.packed_b_layout(depth_block, b.cols(), None);
let tail_layout = if !b.rows().is_multiple_of(depth_block) {
Some(kernel.packed_b_layout(b.rows() % depth_block, b.cols(), None))
} else {
None
};
assert_eq!(layout.size() % layout.align(), 0);
let n_blocks = b.rows() / depth_block;
let total_size =
(n_blocks * layout.size()) + tail_layout.as_ref().map(|l| l.size()).unwrap_or(0);
let mut data = PackingBuffer::new();
let uninit_data = data.alloc_in(alloc, total_size, layout.align());
for (row_block, block_data) in
range_chunks(0..b.rows(), depth_block).zip(uninit_data.chunks_mut(layout.size()))
{
kernel.pack_b_block(block_data, b, row_block, 0..b.cols(), None);
}
unsafe {
data.set_len(total_size);
}
PackedBMatrix {
base: PackedMatrixBase {
data,
depth_size: b.rows(),
nm_size: b.cols(),
panel_size: kernel.nr(),
depth_block,
panel_stride: layout.panel_stride(),
tail_panel_stride: tail_layout
.map(|tl| tl.panel_stride())
.unwrap_or(layout.panel_stride()),
depth_block_stride: layout.size(),
kernel_name: kernel.name(),
},
_marker: PhantomData,
}
}