use std::marker::PhantomData;
use crate::{
ir::{self, Operation},
unexpanded,
};
use super::{
CubeContext, CubePrimitive, CubeType, ExpandElement, ExpandElementTyped, Init, IntoRuntime,
Slice, SliceMut,
};
pub use ir::{MatrixIdent, MatrixLayout};
#[derive(Copy, Clone)]
pub struct Matrix<C: CubeType> {
_c: PhantomData<C>,
}
#[derive(Clone)]
pub struct MatrixExpand {
elem: ExpandElement,
ident: MatrixIdent,
}
impl<C: CubeType> CubeType for Matrix<C> {
type ExpandType = MatrixExpand;
}
impl<C: CubeType> IntoRuntime for Matrix<C> {
fn __expand_runtime_method(self, _context: &mut CubeContext) -> MatrixExpand {
unimplemented!("Matrices can't exist at compile time")
}
}
impl Init for MatrixExpand {
fn init(self, _context: &mut CubeContext) -> Self {
self
}
}
impl<C: CubePrimitive> Matrix<C> {
#[allow(unused_variables)]
pub unsafe fn uninitialized(
ident: MatrixIdent,
m: u32,
n: u32,
k: u32,
layout: MatrixLayout,
) -> Self {
Matrix { _c: PhantomData }
}
#[allow(unused_variables)]
pub fn from_value(
ident: MatrixIdent,
m: u32,
n: u32,
k: u32,
layout: MatrixLayout,
value: C,
) -> Self {
Matrix { _c: PhantomData }
}
#[allow(unused_variables)]
pub fn from_slice(
ident: MatrixIdent,
m: u32,
n: u32,
k: u32,
layout: MatrixLayout,
value: &Slice<'_, C>,
stride: u32,
) -> Self {
Matrix { _c: PhantomData }
}
pub fn __expand_uninitialized(
context: &mut CubeContext,
ident: MatrixIdent,
m: ExpandElementTyped<u32>,
n: ExpandElementTyped<u32>,
k: ExpandElementTyped<u32>,
layout: MatrixLayout,
) -> MatrixExpand {
let elem = context.create_matrix(ir::Matrix {
ident,
m: m.constant().unwrap().as_u32() as u8,
n: n.constant().unwrap().as_u32() as u8,
k: k.constant().unwrap().as_u32() as u8,
elem: C::as_elem(),
layout,
});
MatrixExpand { elem, ident }
}
pub fn __expand_from_value(
context: &mut CubeContext,
ident: MatrixIdent,
m: ExpandElementTyped<u32>,
n: ExpandElementTyped<u32>,
k: ExpandElementTyped<u32>,
layout: MatrixLayout,
value: ExpandElementTyped<C>,
) -> MatrixExpand {
let mat = Self::__expand_uninitialized(context, ident, m, n, k, layout);
fill::expand(context, mat.clone(), value);
mat
}
#[allow(clippy::too_many_arguments)]
pub fn __expand_from_slice(
context: &mut CubeContext,
ident: MatrixIdent,
m: ExpandElementTyped<u32>,
n: ExpandElementTyped<u32>,
k: ExpandElementTyped<u32>,
layout: MatrixLayout,
value: ExpandElementTyped<Slice<'static, C>>,
stride: ExpandElementTyped<u32>,
) -> MatrixExpand {
let mat = Self::__expand_uninitialized(context, ident, m, n, k, layout);
load::expand(context, mat.clone(), value, stride);
mat
}
}
#[allow(unused_variables)]
pub fn fill<C: CubeType>(mat: &Matrix<C>, value: C) {
unexpanded!()
}
pub mod fill {
use super::*;
pub fn expand<C: CubeType>(
context: &mut CubeContext,
mat: MatrixExpand,
value: ExpandElementTyped<C>,
) {
let value: ExpandElement = value.into();
context.register(Operation::CoopMma(ir::CoopMma::Fill {
mat: *mat.elem,
value: *value,
}));
}
}
#[allow(unused_variables)]
pub fn load<C: CubeType>(mat: &Matrix<C>, value: &Slice<'_, C>, stride: u32) {
unexpanded!()
}
pub mod load {
use super::*;
#[allow(unused_variables)]
pub fn expand<C: CubeType>(
context: &mut CubeContext,
mat: MatrixExpand,
value: ExpandElementTyped<Slice<'static, C>>,
stride: ExpandElementTyped<u32>,
) {
let stride: ExpandElement = stride.into();
assert_ne!(
mat.ident,
MatrixIdent::Accumulator,
"Loading accumulator requires explicit layout. Use `load_with_layout` instead."
);
context.register(Operation::CoopMma(ir::CoopMma::Load {
mat: *mat.elem,
value: *value.expand,
stride: *stride,
layout: None,
}));
}
}
#[allow(unused_variables)]
pub fn load_with_layout<C: CubeType>(
mat: &Matrix<C>,
value: &Slice<'_, C>,
stride: u32,
layout: MatrixLayout,
) {
unexpanded!()
}
pub mod load_with_layout {
use super::*;
#[allow(unused_variables)]
pub fn expand<C: CubeType>(
context: &mut CubeContext,
mat: MatrixExpand,
value: ExpandElementTyped<Slice<'static, C>>,
stride: ExpandElementTyped<u32>,
layout: MatrixLayout,
) {
let stride: ExpandElement = stride.into();
context.register(Operation::CoopMma(ir::CoopMma::Load {
mat: *mat.elem,
value: *value.expand,
stride: *stride,
layout: Some(layout),
}));
}
}
#[allow(unused_variables)]
pub fn store<C: CubePrimitive>(
output: &mut SliceMut<'_, C>,
mat: &Matrix<C>,
stride: u32,
layout: MatrixLayout,
) {
unexpanded!()
}
pub mod store {
use super::*;
#[allow(unused_variables)]
pub fn expand<C: CubePrimitive>(
context: &mut CubeContext,
output: ExpandElementTyped<SliceMut<'static, C>>,
mat: MatrixExpand,
stride: ExpandElementTyped<u32>,
layout: MatrixLayout,
) {
let stride: ExpandElement = stride.into();
context.register(Operation::CoopMma(ir::CoopMma::Store {
output: *output.expand,
mat: *mat.elem,
stride: *stride,
layout,
}));
}
}
#[allow(unused_variables)]
pub fn execute<A: CubePrimitive, B: CubePrimitive, C: CubePrimitive, D: CubePrimitive>(
mat_a: &Matrix<A>,
mat_b: &Matrix<B>,
mat_c: &Matrix<C>,
mat_d: &Matrix<D>,
) {
unexpanded!()
}
pub mod execute {
use super::*;
pub fn expand<A: CubePrimitive, B: CubePrimitive, C: CubePrimitive, D: CubePrimitive>(
context: &mut CubeContext,
mat_a: MatrixExpand,
mat_b: MatrixExpand,
mat_c: MatrixExpand,
mat_d: MatrixExpand,
) {
context.register(Operation::CoopMma(ir::CoopMma::Execute {
mat_a: *mat_a.elem,
mat_b: *mat_b.elem,
mat_c: *mat_c.elem,
mat_d: *mat_d.elem,
}));
}
}