use cubecl_core as cubecl;
use cubecl_core::prelude::*;
use serde::{Deserialize, Serialize};
use super::{MatmulIdent, MatmulProblemSize};
#[derive(Clone, Debug)]
pub struct MatmulProblem {
pub m: usize,
pub n: usize,
pub k: usize,
pub lhs_batches: Vec<usize>,
pub rhs_batches: Vec<usize>,
pub out_batches: Vec<usize>,
pub lhs_layout: MatrixLayout,
pub rhs_layout: MatrixLayout,
}
impl MatmulProblem {
fn output_batch_dims(&self) -> Vec<usize> {
self.lhs_batches
.iter()
.rev()
.zip(self.rhs_batches.iter().rev())
.map(|(&dim_lhs, &dim_rhs)| std::cmp::max(dim_lhs, dim_rhs))
.collect()
}
pub(crate) fn num_batches(&self) -> usize {
self.output_batch_dims().iter().product()
}
#[allow(unused)]
pub(crate) fn shape(&self, ident: MatmulIdent) -> Vec<usize> {
match ident {
MatmulIdent::Lhs => self
.lhs_batches
.iter()
.cloned()
.chain(vec![self.m, self.k])
.collect(),
MatmulIdent::Rhs => self
.rhs_batches
.iter()
.cloned()
.chain(vec![self.k, self.n])
.collect(),
MatmulIdent::Out => self
.output_batch_dims()
.iter()
.cloned()
.chain(vec![self.m, self.n])
.collect(),
}
}
}
#[derive(Hash, Eq, PartialEq, Debug, Clone, Serialize, Deserialize)]
pub enum MatmulKind {
General,
MatVec,
VecMat,
ScalarVec,
VecScalar,
InnerProduct,
OuterProduct,
ScalarProduct,
}
impl From<MatmulProblemSize> for MatmulKind {
fn from(matmul_size: MatmulProblemSize) -> Self {
enum DimKind {
Scalar,
Vector,
}
impl From<u32> for DimKind {
fn from(x: u32) -> Self {
match x {
1 => DimKind::Scalar,
_ => DimKind::Vector,
}
}
}
use DimKind::*;
match (
matmul_size.m().into(),
matmul_size.n().into(),
matmul_size.k().into(),
) {
(Scalar, Scalar, Scalar) => MatmulKind::ScalarProduct,
(Scalar, Scalar, Vector) => MatmulKind::InnerProduct,
(Scalar, Vector, Scalar) => MatmulKind::ScalarVec,
(Scalar, Vector, Vector) => MatmulKind::VecMat,
(Vector, Scalar, Scalar) => MatmulKind::VecScalar,
(Vector, Scalar, Vector) => MatmulKind::MatVec,
(Vector, Vector, Scalar) => MatmulKind::OuterProduct,
(Vector, Vector, Vector) => MatmulKind::General,
}
}
}
impl From<MatmulProblem> for MatmulProblemSize {
fn from(problem: MatmulProblem) -> Self {
MatmulProblemSize::new(problem.m as u32, problem.n as u32, problem.k as u32)
}
}
impl From<MatmulProblem> for MatmulKind {
fn from(problem: MatmulProblem) -> Self {
MatmulProblemSize::new(problem.m as u32, problem.n as u32, problem.k as u32).into()
}
}
impl From<&MatmulProblem> for MatmulKind {
fn from(problem: &MatmulProblem) -> Self {
MatmulProblemSize::new(problem.m as u32, problem.n as u32, problem.k as u32).into()
}
}
#[derive(CubeType, Copy, Clone, PartialEq, Eq, Hash, Debug, Default)]
pub enum MatrixLayout {
#[default]
RowMajor,
ColMajor,
}
#[cube]
pub fn as_cmma_layout(#[comptime] layout: MatrixLayout) -> cmma::MatrixLayout {
match layout {
MatrixLayout::RowMajor => cmma::MatrixLayout::RowMajor,
MatrixLayout::ColMajor => cmma::MatrixLayout::ColMajor,
}
}