use std::mem::MaybeUninit;
use std::ops::Range;
use rten_base::byte_cast::{AsBytes, cast_uninit_pod_mut_slice};
use rten_simd::ops::{MaskOps, NumOps};
use rten_simd::{Isa, Mask, Simd};
use rten_tensor::{NdTensorView, Storage};
use super::packing::int8::{PackedBMeta, shift_cast_i8_u8};
pub struct RowOffsets {
pub chan: Vec<i32>,
pub y: Vec<i32>,
pub x: Vec<i32>,
}
pub struct ColOffsets {
pub y: Vec<i32>,
pub x: Vec<i32>,
}
pub struct Im2Col<'a, T> {
pub image: NdTensorView<'a, T, 3>,
pub row_offsets: RowOffsets,
pub col_offsets: ColOffsets,
pub n_cols: usize,
pub n_rows: usize,
pub max_y_offset: i32,
pub max_x_offset: i32,
}
impl<T: Copy + Default> Im2Col<'_, T> {
pub fn rows(&self) -> usize {
self.n_rows
}
pub fn cols(&self) -> usize {
self.n_cols
}
#[inline(always)]
pub(super) fn pack_block<I: Isa, const NR_REGS: usize>(
&self,
isa: I,
out: &mut [MaybeUninit<T>],
panel_width: usize,
rows: Range<usize>,
cols: Range<usize>,
) {
let ops = isa.i32();
let mask_ops = isa.m32();
assert_eq!(panel_width, ops.len() * NR_REGS);
let col_range = cols.start..cols.end.next_multiple_of(panel_width);
let used_size = rows.len() * col_range.len();
assert_eq!(out.len(), used_size);
let col_y_offsets = &self.col_offsets.y[col_range.clone()];
let col_x_offsets = &self.col_offsets.x[col_range.clone()];
let row_chan_offsets = &self.row_offsets.chan[rows.clone()];
let row_y_offsets = &self.row_offsets.y[rows.clone()];
let row_x_offsets = &self.row_offsets.x[rows.clone()];
let img_data = self.image.storage();
let img_len = self.image.storage().len();
assert!(img_len > 0 && img_len <= i32::MAX as usize);
let max_img_offset = ops.splat(img_len as i32 - 1);
let mut out_offset = 0;
for start_col in (0..col_y_offsets.len()).step_by(ops.len() * NR_REGS) {
let col_y_offset: [I::I32; NR_REGS] =
std::array::from_fn(|i| ops.load(&col_y_offsets[start_col + ops.len() * i..]));
let col_x_offset: [I::I32; NR_REGS] =
std::array::from_fn(|i| ops.load(&col_x_offsets[start_col + ops.len() * i..]));
let max_x_offset = ops.splat(self.max_x_offset);
let max_y_offset = ops.splat(self.max_y_offset);
for ((&row_chan_offset, &row_y_offset), &row_x_offset) in row_chan_offsets
.iter()
.zip(row_y_offsets.iter())
.zip(row_x_offsets.iter())
{
let row_chan_offset = ops.splat(row_chan_offset);
let row_y_offset = ops.splat(row_y_offset);
let row_x_offset = ops.splat(row_x_offset);
for i in 0..NR_REGS {
let y_offset = ops.add(col_y_offset[i], row_y_offset);
let x_offset = ops.add(col_x_offset[i], row_x_offset);
let offsets = ops.add(ops.add(row_chan_offset, y_offset), x_offset);
let offsets = ops.min(ops.max(offsets, ops.zero()), max_img_offset);
let zero = ops.zero();
let y_valid =
mask_ops.and(ops.ge(y_offset, zero), ops.le(y_offset, max_y_offset));
let x_valid =
mask_ops.and(ops.ge(x_offset, zero), ops.le(x_offset, max_x_offset));
let pad_mask = mask_ops.and(y_valid, x_valid);
let offsets_array = ops.select(offsets, zero, pad_mask).to_array();
let pad_mask_array = pad_mask.to_array();
for idx in 0..ops.len() {
let src_elem =
unsafe { *img_data.get_unchecked(offsets_array[idx] as usize) };
let elem = if pad_mask_array[idx] {
src_elem
} else {
T::default()
};
let out_el = unsafe { out.get_unchecked_mut(out_offset + idx) };
out_el.write(elem);
}
out_offset += ops.len();
}
}
}
assert_eq!(out_offset, used_size);
}
}
impl Im2Col<'_, i8> {
#[inline(always)]
#[allow(unused)] pub(super) fn pack_block_i8_dot<
I: Isa,
const NR: usize,
const NR_REGS: usize,
const K_TILE: usize,
>(
&self,
isa: I,
out: &mut [MaybeUninit<i8>],
rows: Range<usize>,
cols: Range<usize>,
zero_point: i8,
) {
self.pack_block_int8::<_, NR, NR_REGS, K_TILE, false>(isa, out, rows, cols, zero_point);
}
#[inline(always)]
#[allow(unused)] pub(super) fn pack_block_i8_dot_cast_u8<
I: Isa,
const NR: usize,
const NR_REGS: usize,
const K_TILE: usize,
>(
&self,
isa: I,
out: &mut [MaybeUninit<u8>],
rows: Range<usize>,
cols: Range<usize>,
zero_point: i8,
) {
let out = cast_uninit_pod_mut_slice(out).unwrap();
self.pack_block_int8::<_, NR, NR_REGS, K_TILE, true>(isa, out, rows, cols, zero_point);
}
#[inline(always)]
fn pack_block_int8<
I: Isa,
const NR: usize,
const NR_REGS: usize,
const K_TILE: usize,
const CAST_B_U8: bool,
>(
&self,
isa: I,
out: &mut [MaybeUninit<i8>],
rows: Range<usize>,
cols: Range<usize>,
zero_point: i8,
) {
let ops = isa.i32();
assert_eq!(ops.len() * NR_REGS, NR);
let mask_ops = isa.m32();
debug_assert!(rows.end <= self.rows());
debug_assert!(cols.end <= self.cols());
let max_x_offset = ops.splat(self.max_x_offset);
let max_y_offset = ops.splat(self.max_y_offset);
let col_x_offsets = &self.col_offsets.x;
debug_assert_eq!(col_x_offsets.len() % ops.len(), 0);
let col_y_offsets = &self.col_offsets.y;
debug_assert_eq!(col_y_offsets.len() % ops.len(), 0);
let row_x_offsets = &self.row_offsets.x;
debug_assert_eq!(row_x_offsets.len() % K_TILE, 0);
let row_y_offsets = &self.row_offsets.y;
debug_assert_eq!(row_y_offsets.len() % K_TILE, 0);
let row_chan_offsets = &self.row_offsets.chan;
debug_assert_eq!(row_chan_offsets.len() % K_TILE, 0);
let img_data = self.image.storage();
let mut out_offset = 0;
for start_col in cols.step_by(ops.len() * NR_REGS) {
let col_y_offset: [I::I32; NR_REGS] =
std::array::from_fn(|i| ops.load(&col_y_offsets[start_col + i * ops.len()..]));
let col_x_offset: [I::I32; NR_REGS] =
std::array::from_fn(|i| ops.load(&col_x_offsets[start_col + i * ops.len()..]));
let zero = ops.zero();
let mut col_sums = [ops.zero().to_array(); NR_REGS];
for start_row in rows.clone().step_by(K_TILE) {
for i in 0..K_TILE {
let k = start_row + i;
let row_x_offset = ops.splat(unsafe { *row_x_offsets.get_unchecked(k) });
let row_y_offset = ops.splat(unsafe { *row_y_offsets.get_unchecked(k) });
let row_chan_offset = ops.splat(unsafe { *row_chan_offsets.get_unchecked(k) });
for c_block in 0..NR_REGS {
let x_offsets = ops.add(row_x_offset, col_x_offset[c_block]);
let y_offsets = ops.add(row_y_offset, col_y_offset[c_block]);
let offsets = ops.add(ops.add(x_offsets, y_offsets), row_chan_offset);
let y_valid =
mask_ops.and(ops.ge(y_offsets, zero), ops.le(y_offsets, max_y_offset));
let x_valid =
mask_ops.and(ops.ge(x_offsets, zero), ops.le(x_offsets, max_x_offset));
let pad_mask = mask_ops.and(y_valid, x_valid);
let pad_mask_array = pad_mask.to_array();
let offsets_array = ops.select(offsets, zero, pad_mask).to_array();
for idx in 0..ops.len() {
let out_elem = unsafe {
out.get_unchecked_mut(
out_offset + (c_block * ops.len() + idx) * K_TILE + i,
)
};
let src_elem =
unsafe { *img_data.get_unchecked(offsets_array[idx] as usize) };
if CAST_B_U8 {
let src_elem = shift_cast_i8_u8(src_elem);
let elem = if pad_mask_array[idx] { src_elem } else { 0 };
col_sums[c_block][idx] += elem as i32;
out_elem.write(elem as i8);
} else {
let elem = if pad_mask_array[idx] { src_elem } else { 0 };
col_sums[c_block][idx] += elem as i32;
out_elem.write(elem);
}
}
}
}
out_offset += ops.len() * NR_REGS * K_TILE;
}
let meta = PackedBMeta::<NR> {
col_sums: *unsafe {
std::mem::transmute::<&[<I::I32 as Simd>::Array; NR_REGS], &[i32; NR]>(
&col_sums,
)
},
zero_points: if CAST_B_U8 {
[shift_cast_i8_u8(zero_point) as i32; NR]
} else {
[zero_point as i32; NR]
},
};
let meta_bytes = meta.as_bytes();
for (i, byte) in meta_bytes.iter().enumerate() {
out[out_offset + i].write(*byte as i8);
}
out_offset += meta_bytes.len();
}
assert_eq!(out_offset, out.len());
}
}