use cubecl_core as cubecl;
use cubecl_core::prelude::*;
use crate::components::{MatrixLayout, stage::StageMemoryConfig};
#[derive(CubeType, Clone, Copy)]
pub struct StridedTile<ES: Numeric, IO: SliceVisibility = ReadOnly> {
pub slice: Slice<Line<ES>, IO>,
pub stride: u32,
#[cube(comptime)]
pub layout: MatrixLayout,
}
#[cube]
impl<ES: Numeric> StridedTile<ES> {
pub fn new_contiguous(
slice: Slice<Line<ES>>,
#[comptime] config: StageMemoryConfig,
) -> StridedTile<ES> {
let layout = config.matrix_layout;
let stride = match layout {
MatrixLayout::RowMajor => config.elements_in_tile_col,
MatrixLayout::ColMajor => config.elements_in_tile_row,
};
let stride = comptime![stride / config.stage_line_size];
StridedTile::<ES> {
slice,
stride,
layout,
}
}
pub fn new_contiguous_mut(
slice: Slice<Line<ES>, ReadWrite>,
#[comptime] config: StageMemoryConfig,
) -> StridedTile<ES, ReadWrite> {
let layout = config.matrix_layout;
let stride = match layout {
MatrixLayout::RowMajor => config.elements_in_tile_col,
MatrixLayout::ColMajor => config.elements_in_tile_row,
};
let stride = comptime![stride / config.stage_line_size];
StridedTile::<ES, ReadWrite> {
slice,
stride,
layout,
}
}
pub fn new_strided(
slice: Slice<Line<ES>>,
stride: u32,
#[comptime] layout: MatrixLayout,
) -> StridedTile<ES> {
StridedTile::<ES> {
slice,
stride,
layout,
}
}
pub fn new_strided_mut(
slice: Slice<Line<ES>, ReadWrite>,
stride: u32,
#[comptime] layout: MatrixLayout,
) -> StridedTile<ES, ReadWrite> {
StridedTile::<ES, ReadWrite> {
slice,
stride,
layout,
}
}
}
#[cube]
impl<ES: Numeric, IO: SliceVisibility> StridedTile<ES, IO> {
pub fn as_unlined(&self) -> (Slice<ES, IO>, u32) {
let stage_line_size = comptime![self.slice.line_size()];
(
self.slice.try_cast_unchecked(),
self.stride * stage_line_size,
)
}
pub fn get_line(&self, coor_strided: u32, coor_contiguous: u32) -> Line<ES> {
self.slice[coor_strided * self.stride + coor_contiguous]
}
}