use std::mem::MaybeUninit;
use rten_base::byte_cast::{AsBytes, FromBytes};
use rten_tensor::layout::AsIndex;
use rten_tensor::prelude::*;
use rten_tensor::storage::ViewData;
use rten_tensor::{Layout, Matrix, NdIndices, NdLayout, TensorBase};
use super::PackedLayout;
use super::SliceWriter;
#[derive(Copy, Clone)]
#[repr(C)]
pub struct PackedAMeta<const MR: usize> {
pub row_sums: [i32; MR],
pub zero_points: [i32; MR],
}
unsafe impl<const MR: usize> AsBytes for PackedAMeta<MR> {}
unsafe impl<const MR: usize> FromBytes for PackedAMeta<MR> {}
#[derive(Copy, Clone)]
#[repr(C)]
pub struct PackedBMeta<const NR: usize> {
pub col_sums: [i32; NR],
pub zero_points: [i32; NR],
}
unsafe impl<const NR: usize> AsBytes for PackedBMeta<NR> {}
unsafe impl<const NR: usize> FromBytes for PackedBMeta<NR> {}
pub fn packed_b_layout<const NR: usize, const K_TILE: usize>(
b_rows: usize,
b_cols: usize,
) -> PackedLayout {
let n_panels = b_cols.div_ceil(NR);
let packed_elements_size = b_rows.div_ceil(K_TILE) * NR * K_TILE;
debug_assert_eq!(packed_elements_size % align_of::<PackedBMeta<NR>>(), 0);
let panel_stride = packed_elements_size + size_of::<PackedBMeta<NR>>();
let size = n_panels * panel_stride;
let align = align_of::<PackedBMeta<NR>>();
PackedLayout::new(size, align, panel_stride)
}
#[allow(unused)]
pub fn pack_b<const NR: usize, const K_TILE: usize>(
out: &mut [MaybeUninit<i8>],
b: Matrix<i8>,
zero_points: Option<&[i8]>,
) {
pack_b_impl::<NR, K_TILE, _>(
out,
b,
zero_points,
|x| x,
|out, meta| {
out.write_slice(meta.as_signed_bytes());
},
)
}
#[inline]
pub fn shift_cast_i8_u8(x: i8) -> u8 {
x as u8 ^ 0x80
}
#[allow(unused)]
pub fn pack_b_cast_i8_u8<const NR: usize, const K_TILE: usize>(
out: &mut [MaybeUninit<u8>],
b: Matrix<i8>,
zero_points: Option<&[i8]>,
) {
pack_b_impl::<NR, K_TILE, _>(out, b, zero_points, shift_cast_i8_u8, |out, meta| {
out.write_slice(meta.as_bytes())
})
}
trait Byte: Copy + Default {}
impl Byte for u8 {}
impl Byte for i8 {}
fn pack_b_impl<const NR: usize, const K_TILE: usize, T: Byte>(
out: &mut [MaybeUninit<T>],
b: Matrix<i8>,
zero_point: Option<&[i8]>,
cast: impl Fn(i8) -> T,
write_meta: impl Fn(&mut SliceWriter<T>, PackedBMeta<NR>),
) where
i32: From<T>,
{
let [b_rows, b_cols] = b.shape();
assert_eq!(
out.len(),
packed_b_layout::<NR, K_TILE>(b_rows, b_cols).size()
);
let mut out = SliceWriter::new(out);
let full_panels = b_cols / NR;
let tail_cols = b_cols % NR;
let full_tiles = b_rows / K_TILE;
let tail_rows = b_rows % K_TILE;
for col_panel in 0..full_panels {
let mut col_sums = [0i32; NR];
for row_tile in 0..full_tiles {
for c in 0..NR {
for r in 0..K_TILE {
let y = row_tile * K_TILE + r;
let x = col_panel * NR + c;
let val = cast(unsafe { *b.get_unchecked([y, x]) });
col_sums[c] += i32::from(val);
unsafe { out.write_unchecked(val) };
}
}
}
if tail_rows != 0 {
for c in 0..NR {
for r in 0..tail_rows {
let y = full_tiles * K_TILE + r;
let x = col_panel * NR + c;
unsafe {
let val = cast(*b.get_unchecked([y, x]));
col_sums[c] += i32::from(val);
out.write_unchecked(val);
}
}
unsafe { out.write_n_unchecked(K_TILE - tail_rows, T::default()) };
}
}
let meta = PackedBMeta {
col_sums,
zero_points: if let Some(zp) = zero_point {
std::array::from_fn(|c| i32::from(cast(zp[c])))
} else {
[i32::from(cast(0)); NR]
},
};
write_meta(&mut out, meta);
}
if tail_cols != 0 {
let mut col_sums = [0i32; NR];
let col_panel = full_panels;
for row_tile in 0..full_tiles {
for c in 0..tail_cols {
for r in 0..K_TILE {
let y = row_tile * K_TILE + r;
let x = col_panel * NR + c;
let val = cast(unsafe { *b.get_unchecked([y, x]) });
col_sums[c] += i32::from(val);
unsafe { out.write_unchecked(val) };
}
}
unsafe { out.write_n_unchecked((NR - tail_cols) * K_TILE, T::default()) };
}
if tail_rows != 0 {
for c in 0..tail_cols {
for r in 0..tail_rows {
let y = full_tiles * K_TILE + r;
let x = col_panel * NR + c;
unsafe {
let val = cast(*b.get_unchecked([y, x]));
col_sums[c] += i32::from(val);
out.write_unchecked(val);
}
}
unsafe { out.write_n_unchecked(K_TILE - tail_rows, T::default()) };
}
unsafe { out.write_n_unchecked((NR - tail_cols) * K_TILE, T::default()) };
}
let mut zero_point_array = [0i32; NR];
if let Some(zp) = zero_point {
let col_range = col_panel * NR..b_cols;
for (i, c) in col_range.enumerate() {
zero_point_array[i] = i32::from(cast(zp[c]));
}
} else {
for i in 0..tail_cols {
zero_point_array[i] = i32::from(cast(0));
}
}
let meta = PackedBMeta {
col_sums,
zero_points: zero_point_array,
};
write_meta(&mut out, meta);
}
assert!(out.completed());
}
pub fn packed_a_layout<const MR: usize, const K_TILE: usize>(
a_rows: usize,
a_cols: usize,
) -> PackedLayout {
let n_panels = a_rows.div_ceil(MR);
let packed_elements_size = a_cols.div_ceil(K_TILE) * MR * K_TILE;
debug_assert_eq!(packed_elements_size % align_of::<PackedAMeta<MR>>(), 0);
let panel_stride = packed_elements_size + size_of::<PackedAMeta<MR>>();
let size = n_panels * panel_stride;
let align = align_of::<PackedAMeta<MR>>();
PackedLayout::new(size, align, panel_stride)
}
#[derive(Clone)]
struct RowMajorLayout {
shape: [usize; 2],
row_stride: usize,
}
impl RowMajorLayout {
fn from_layout(layout: NdLayout<2>) -> Option<Self> {
if layout.stride(1) == 1 {
Some(Self {
shape: layout.shape(),
row_stride: layout.stride(0),
})
} else {
None
}
}
fn index_valid(&self, index: [usize; 2]) -> bool {
index[0] < self.shape[0] && index[1] < self.shape[1]
}
}
impl Layout for RowMajorLayout {
type Index<'a> = [usize; 2];
type Indices = NdIndices<2>;
fn ndim(&self) -> usize {
2
}
fn len(&self) -> usize {
self.shape.iter().product()
}
#[inline]
fn offset(&self, index: [usize; 2]) -> Option<usize> {
self.index_valid(index)
.then_some(self.offset_unchecked(index))
}
#[inline]
fn offset_unchecked(&self, index: [usize; 2]) -> usize {
index[0] * self.row_stride + index[1]
}
#[inline]
fn shape(&self) -> Self::Index<'_> {
self.shape
}
#[inline]
fn strides(&self) -> Self::Index<'_> {
[self.row_stride, 1]
}
fn indices(&self) -> Self::Indices {
NdIndices::from_shape(self.shape)
}
}
impl AsIndex<RowMajorLayout> for [usize; 2] {
fn as_index(&self) -> [usize; 2] {
*self
}
}
pub fn pack_a<const MR: usize, const K_TILE: usize>(
out: &mut [MaybeUninit<u8>],
a: Matrix<u8>,
zero_point: Option<&[u8]>,
) {
if let Some(layout) = RowMajorLayout::from_layout(*a.layout()) {
pack_a_impl::<MR, K_TILE, _>(
out,
TensorBase::from_storage_and_layout(a.storage(), layout),
zero_point,
)
} else {
pack_a_impl::<MR, K_TILE, _>(out, a, zero_point)
}
}
#[inline(never)]
fn pack_a_impl<const MR: usize, const K_TILE: usize, L>(
out: &mut [MaybeUninit<u8>],
a: TensorBase<ViewData<u8>, L>,
zero_point: Option<&[u8]>,
) where
L: Clone + for<'a> Layout<Index<'a> = [usize; 2]>,
[usize; 2]: AsIndex<L>,
{
let [a_rows, a_cols] = a.shape();
assert_eq!(
out.len(),
packed_a_layout::<MR, K_TILE>(a_rows, a_cols).size()
);
let mut out = SliceWriter::new(out);
let full_panels = a_rows / MR;
let tail_rows = a_rows % MR;
let full_tiles = a_cols / K_TILE;
let tail_cols = a_cols % K_TILE;
for row_tile in 0..full_panels {
let mut row_sums = [0i32; MR];
for col_tile in 0..full_tiles {
for r in 0..MR {
for c in 0..K_TILE {
let y = row_tile * MR + r;
let x = col_tile * K_TILE + c;
let val = unsafe { *a.get_unchecked([y, x]) };
row_sums[r] += val as i32;
unsafe { out.write_unchecked(val) };
}
}
}
if tail_cols != 0 {
for r in 0..MR {
for c in 0..tail_cols {
let y = row_tile * MR + r;
let x = full_tiles * K_TILE + c;
let val = unsafe { *a.get_unchecked([y, x]) };
row_sums[r] += val as i32;
unsafe { out.write_unchecked(val) };
}
unsafe { out.write_n_unchecked(K_TILE - tail_cols, 0) };
}
}
let meta = PackedAMeta {
row_sums,
zero_points: if let Some(zp) = zero_point {
std::array::from_fn(|r| zp[r] as i32)
} else {
[0; MR]
},
};
out.write_slice(meta.as_bytes());
}
if tail_rows != 0 {
let row_tile = full_panels;
let mut row_sums = [0i32; MR];
let row_range = row_tile * MR..(row_tile * MR + MR).min(a_rows);
for col_tile in 0..full_tiles {
for r in 0..tail_rows {
for c in 0..K_TILE {
let y = row_tile * MR + r;
let x = col_tile * K_TILE + c;
let val = unsafe { *a.get_unchecked([y, x]) };
row_sums[r] += val as i32;
unsafe { out.write_unchecked(val) };
}
}
unsafe { out.write_n_unchecked((MR - tail_rows) * K_TILE, 0) };
}
if tail_cols != 0 {
for r in 0..tail_rows {
for c in 0..tail_cols {
let y = row_tile * MR + r;
let x = full_tiles * K_TILE + c;
let val = unsafe { *a.get_unchecked([y, x]) };
row_sums[r] += val as i32;
unsafe { out.write_unchecked(val) };
}
unsafe { out.write_n_unchecked(K_TILE - tail_cols, 0) };
}
unsafe { out.write_n_unchecked((MR - tail_rows) * K_TILE, 0) };
}
let mut zero_point_array = [0i32; MR];
if let Some(zp) = zero_point {
for (i, r) in row_range.enumerate() {
zero_point_array[i] = zp[r] as i32;
}
}
let meta = PackedAMeta {
row_sums,
zero_points: zero_point_array,
};
out.write_slice(meta.as_bytes());
}
assert!(out.completed());
}
pub fn extract_packed_a<const MR: usize>(a: &[u8]) -> (&[u8], &PackedAMeta<MR>) {
assert!(a.len() >= size_of::<PackedAMeta<MR>>());
let meta_offset = a.len() - size_of::<PackedAMeta<MR>>();
let (packed_elements, meta_bytes) = a.split_at(meta_offset);
(packed_elements, PackedAMeta::from_bytes(meta_bytes))
}
pub fn extract_packed_b<const NR: usize>(b: &[u8]) -> (&[u8], &PackedBMeta<NR>) {
assert!(b.len() >= size_of::<PackedBMeta<NR>>());
let meta_offset = b.len() - size_of::<PackedBMeta<NR>>();
let (packed_elements, meta_bytes) = b.split_at(meta_offset);
(packed_elements, PackedBMeta::from_bytes(meta_bytes))
}
#[cfg(test)]
mod tests {
use rten_base::byte_cast::{AsBytes, cast_pod_slice};
use rten_tensor::prelude::*;
use rten_tensor::rng::XorShiftRng;
use rten_tensor::{Matrix, MatrixLayout, NdTensor};
use super::{
PackedAMeta, PackedBMeta, extract_packed_a, extract_packed_b, pack_a, pack_b,
pack_b_cast_i8_u8, packed_a_layout, packed_b_layout,
};
const K_TILE_I8DOT: usize = 4;
const K_TILE_I8MM: usize = 8;
fn pack_a_matrix<const MR: usize, const K_TILE: usize>(mat: Matrix<u8>) -> Vec<u8> {
let layout = packed_a_layout::<MR, K_TILE>(mat.rows(), mat.cols());
assert!(layout.size() >= mat.rows() * mat.cols() + 2 * mat.rows() * 4);
let mut buf = Vec::with_capacity(layout.size());
pack_a::<MR, K_TILE>(
&mut buf.spare_capacity_mut()[..layout.size()],
mat.view(),
None,
);
unsafe { buf.set_len(layout.size()) }
buf
}
fn reference_pack_a<const MR: usize, const K_TILE: usize>(mat: Matrix<u8>) -> Vec<u8> {
let layout = packed_a_layout::<MR, K_TILE>(mat.rows(), mat.cols());
let mut buf = Vec::with_capacity(layout.size());
for row_panel in 0..mat.rows().div_ceil(MR) {
let mut row_sums = [0i32; MR];
for k_tile in 0..mat.cols().div_ceil(K_TILE) {
for r in 0..MR {
for c in 0..K_TILE {
let y = row_panel * MR + r;
let x = k_tile * K_TILE + c;
let val = mat.get([y, x]).copied().unwrap_or(0);
row_sums[r] += val as i32;
buf.push(val);
}
}
}
let meta = PackedAMeta {
row_sums,
zero_points: [0; MR],
};
buf.extend_from_slice(meta.as_bytes());
}
assert_eq!(buf.len(), layout.size());
buf
}
fn pack_b_matrix<const NR: usize, const K_TILE: usize>(mat: Matrix<i8>) -> Vec<i8> {
let layout = packed_b_layout::<NR, K_TILE>(mat.rows(), mat.cols());
assert!(layout.size() >= mat.rows() * mat.cols() + mat.cols() * 4);
let mut buf = Vec::with_capacity(layout.size());
pack_b::<NR, K_TILE>(
&mut buf.spare_capacity_mut()[..layout.size()],
mat.view(),
None,
);
unsafe { buf.set_len(layout.size()) }
buf
}
fn reference_pack_b<const NR: usize, const K_TILE: usize>(mat: Matrix<i8>) -> Vec<i8> {
let layout = packed_b_layout::<NR, K_TILE>(mat.rows(), mat.cols());
let mut buf = Vec::with_capacity(layout.size());
for col_panel in 0..mat.cols().div_ceil(NR) {
let mut col_sums = [0i32; NR];
for k_tile in 0..mat.rows().div_ceil(K_TILE) {
for c in 0..NR {
for r in 0..K_TILE {
let y = k_tile * K_TILE + r;
let x = col_panel * NR + c;
let val = mat.get([y, x]).copied().unwrap_or(0);
col_sums[c] += val as i32;
buf.push(val);
}
}
}
let meta = PackedBMeta {
col_sums,
zero_points: [0; NR],
};
buf.extend_from_slice(cast_pod_slice(meta.as_bytes()).unwrap());
}
assert_eq!(buf.len(), layout.size());
buf
}
fn pack_b_matrix_cast_u8<const NR: usize, const K_TILE: usize>(mat: Matrix<i8>) -> Vec<u8> {
let layout = packed_b_layout::<NR, K_TILE>(mat.rows(), mat.cols());
assert!(layout.size() >= mat.rows() * mat.cols() + mat.cols() * 4);
let mut buf = Vec::with_capacity(layout.size());
pack_b_cast_i8_u8::<NR, K_TILE>(
&mut buf.spare_capacity_mut()[..layout.size()],
mat.view(),
None,
);
unsafe { buf.set_len(layout.size()) }
buf
}
#[test]
fn test_pack_a_various_sizes() {
const MR: usize = 8;
fn test_pack_a<const K_TILE: usize>() {
let mut rng = XorShiftRng::new(5678);
for m in 1..MR * 2 {
for k in 1..K_TILE_I8DOT * 2 {
let mat = NdTensor::rand([m, k], &mut rng);
let expected = reference_pack_a::<MR, K_TILE_I8DOT>(mat.view());
let actual = pack_a_matrix::<MR, K_TILE_I8DOT>(mat.view());
assert_eq!(
actual, expected,
"packed buffer mismatch for row-major m={} k={}",
m, k
);
let mat = NdTensor::rand([k, m], &mut rng);
let expected = reference_pack_a::<MR, K_TILE_I8DOT>(mat.transposed().view());
let actual = pack_a_matrix::<MR, K_TILE_I8DOT>(mat.transposed().view());
assert_eq!(
actual, expected,
"packed buffer mismatch for col-major m={} k={}",
m, k
);
}
}
}
test_pack_a::<K_TILE_I8DOT>();
test_pack_a::<K_TILE_I8MM>();
}
#[test]
fn test_extract_packed_a() {
fn test_extract_packed_a<const K_TILE: usize>() {
const MR: usize = 8;
let mat = NdTensor::<u8, 2>::from([[1, 2], [3, 4]]);
let packed = pack_a_matrix::<MR, K_TILE_I8DOT>(mat.view());
let (packed_elems, meta) = extract_packed_a(&packed);
assert!(packed_elems.len() >= mat.rows() * mat.cols());
assert_eq!(meta.row_sums, [3, 7, 0, 0, 0, 0, 0, 0]);
}
test_extract_packed_a::<K_TILE_I8DOT>();
test_extract_packed_a::<K_TILE_I8MM>();
}
#[test]
fn test_pack_b_various_sizes() {
const NR: usize = 8;
fn test_pack_b<const K_TILE: usize>() {
let mut rng = XorShiftRng::new(5678);
for n in 1..NR * 2 {
for k in 1..K_TILE * 2 {
let mat = NdTensor::rand([k, n], &mut rng);
let expected = reference_pack_b::<NR, K_TILE>(mat.view());
let actual = pack_b_matrix::<NR, K_TILE>(mat.view());
assert_eq!(
actual, expected,
"packed buffer mismatch for n={} k={}",
n, k
);
}
}
}
test_pack_b::<K_TILE_I8DOT>();
test_pack_b::<K_TILE_I8MM>();
}
#[test]
fn test_extract_packed_b() {
const NR: usize = 8;
let mat = NdTensor::<i8, 2>::from([[1, 2], [3, 4]]);
let packed = pack_b_matrix::<NR, K_TILE_I8DOT>(mat.view());
let (packed_elems, meta) = extract_packed_b(cast_pod_slice(&packed).unwrap());
assert!(packed_elems.len() >= mat.rows() * mat.cols());
assert_eq!(meta.col_sums, [4, 6, 0, 0, 0, 0, 0, 0]);
}
#[test]
fn test_pack_b_cast_i8_u8() {
let mat = NdTensor::<i8, 2>::from([[1, 2], [3, 4]]);
let packed = pack_b_matrix_cast_u8::<2, K_TILE_I8DOT>(mat.view());
let (packed_elems, meta) = extract_packed_b(&packed);
assert!(packed_elems.len() >= mat.rows() * mat.cols());
assert_eq!(packed_elems, &[129, 131, 0, 0, 130, 132, 0, 0]);
assert_eq!(meta.col_sums, [129 + 131, 130 + 132]);
}
}