use crate::align::{Alignment, Unaligned};
use crate::scalar::NumericElement;
use core::marker::PhantomData;
#[repr(C)]
pub struct TileView<
'a,
T: NumericElement,
Backend,
Arch,
const ROWS: usize,
const COLS: usize,
Align: Alignment = Unaligned,
Ref: 'a = &'a [T],
> {
ptr: *mut T,
stride: usize,
_marker: PhantomData<(&'a T, Backend, Arch, Align, Ref)>,
}
unsafe impl<
'a,
T: NumericElement,
Backend,
Arch,
const ROWS: usize,
const COLS: usize,
Align: Alignment,
Ref: 'a,
> Send for TileView<'a, T, Backend, Arch, ROWS, COLS, Align, Ref>
where
Ref: Send,
{
}
unsafe impl<
'a,
T: NumericElement,
Backend,
Arch,
const ROWS: usize,
const COLS: usize,
Align: Alignment,
Ref: 'a,
> Sync for TileView<'a, T, Backend, Arch, ROWS, COLS, Align, Ref>
where
Ref: Sync,
{
}
impl<
'a,
T: NumericElement,
Backend,
Arch,
const ROWS: usize,
const COLS: usize,
Align: Alignment,
Ref: 'a,
> Clone for TileView<'a, T, Backend, Arch, ROWS, COLS, Align, Ref>
where
Ref: Clone,
{
#[inline(always)]
fn clone(&self) -> Self {
Self {
ptr: self.ptr,
stride: self.stride,
_marker: PhantomData,
}
}
}
impl<
'a,
T: NumericElement,
Backend,
Arch,
const ROWS: usize,
const COLS: usize,
Align: Alignment,
Ref: 'a,
> Copy for TileView<'a, T, Backend, Arch, ROWS, COLS, Align, Ref>
where
Ref: Copy,
{
}
impl<
'a,
T: NumericElement,
Backend,
Arch: crate::arch::SimdArch,
const ROWS: usize,
const COLS: usize,
Align: Alignment,
> TileView<'a, T, Backend, Arch, ROWS, COLS, Align, &'a [T]>
{
#[inline]
pub fn new(data: &'a [T], stride: usize) -> Option<Self> {
if data.len() < ROWS * stride {
return None;
}
if Align::IS_ALIGNED {
let req_align = Arch::REGISTER_WIDTH_BITS as usize / 8;
if req_align > 0 && Align::ALIGN_BYTES < req_align {
return None;
}
let addr = data.as_ptr() as usize;
if addr % Align::ALIGN_BYTES != 0 {
return None;
}
}
Some(Self {
ptr: data.as_ptr() as *mut T,
stride,
_marker: PhantomData,
})
}
}
impl<
'a,
T: NumericElement,
Backend,
Arch: crate::arch::SimdArch,
const ROWS: usize,
const COLS: usize,
Align: Alignment,
> TileView<'a, T, Backend, Arch, ROWS, COLS, Align, &'a mut [T]>
{
#[inline]
pub fn new_mut(data: &'a mut [T], stride: usize) -> Option<Self> {
if data.len() < ROWS * stride {
return None;
}
if Align::IS_ALIGNED {
let req_align = Arch::REGISTER_WIDTH_BITS as usize / 8;
if req_align > 0 && Align::ALIGN_BYTES < req_align {
return None;
}
let addr = data.as_ptr() as usize;
if addr % Align::ALIGN_BYTES != 0 {
return None;
}
}
Some(Self {
ptr: data.as_mut_ptr(),
stride,
_marker: PhantomData,
})
}
#[inline(always)]
pub fn as_mut_ptr(&mut self) -> *mut T {
self.ptr
}
}
impl<
'a,
T: NumericElement,
Backend,
Arch,
const ROWS: usize,
const COLS: usize,
Align: Alignment,
Ref: 'a,
> TileView<'a, T, Backend, Arch, ROWS, COLS, Align, Ref>
{
#[inline(always)]
pub fn as_ptr(&self) -> *const T {
self.ptr
}
#[inline(always)]
pub fn stride(&self) -> usize {
self.stride
}
#[inline(always)]
pub fn rows(&self) -> usize {
ROWS
}
#[inline(always)]
pub fn cols(&self) -> usize {
COLS
}
}
pub trait TileMatrixMultiply<
TA,
TB,
TC,
Backend,
Arch,
const M: usize,
const N: usize,
const K: usize,
>
{
unsafe fn tile_matmul(
c: *mut TC,
c_stride: usize,
a: *const TA,
a_stride: usize,
b: *const TB,
b_stride: usize,
);
}