use std::marker::PhantomData;
use rten_tensor::prelude::*;
use rten_tensor::{MatrixLayout, MatrixMut, StorageMut};
pub struct OutputTiles<'a, T> {
data: *mut T,
rows: usize,
cols: usize,
row_stride: usize,
tile_rows: usize,
tile_cols: usize,
n_row_tiles: usize,
n_col_tiles: usize,
_marker: PhantomData<&'a mut [T]>,
}
unsafe impl<T> Sync for OutputTiles<'_, T> {}
impl<'a, T> OutputTiles<'a, T> {
pub fn new(
mut data: MatrixMut<'a, T>,
tile_rows: usize,
tile_cols: usize,
) -> OutputTiles<'a, T> {
OutputTiles {
data: data.storage_mut().as_mut_ptr(),
rows: data.rows(),
cols: data.cols(),
row_stride: data.stride(0),
tile_rows,
tile_cols,
n_row_tiles: data.rows().div_ceil(tile_rows),
n_col_tiles: data.cols().div_ceil(tile_cols),
_marker: PhantomData,
}
}
pub unsafe fn tile(&self, row: usize, col: usize) -> OutputTile<'_, T> {
assert!(row < self.n_row_tiles && col < self.n_col_tiles);
let start_row = row * self.tile_rows;
let start_col = col * self.tile_cols;
OutputTile {
ptr: unsafe { self.data.add(start_row * self.row_stride + start_col) },
row_stride: self.row_stride,
used_rows: (self.rows - start_row).min(self.tile_rows),
used_cols: (self.cols - start_col).min(self.tile_cols),
_marker: PhantomData,
}
}
}
pub struct OutputTile<'a, T> {
pub ptr: *mut T,
pub row_stride: usize,
pub used_rows: usize,
pub used_cols: usize,
_marker: PhantomData<&'a mut [T]>,
}