use crate::{client::ComputeClient, runtime::Runtime};
use cubecl_ir::{ElemType, StorageType};
#[derive(Eq, PartialEq, Clone, Hash, Debug, Copy)]
#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
pub struct ComputeCmmaConfig {
pub accumulator_type: AccumulatorType,
pub cmma_dims: CmmaDims,
}
pub type AccumulatorType = ElemType;
#[derive(Eq, PartialEq, Clone, Hash, Debug, Copy)]
#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
pub struct CmmaDims {
pub m: usize,
pub n: usize,
pub k: usize,
}
impl CmmaDims {
pub fn num_elems(&self) -> usize {
self.m * self.n * self.k
}
}
pub fn select_cmma_tile<R: Runtime>(
client: &ComputeClient<R>,
lhs: StorageType,
rhs: StorageType,
acc: StorageType,
(m, n, k): (usize, usize, usize),
) -> Option<(u32, u32, u32)> {
let props = client.properties();
props
.features
.matmul
.cmma
.iter()
.chain(props.features.matmul.mma.iter())
.filter(|it| it.a_type == lhs && it.b_type == rhs && it.cd_type == acc)
.filter(|it| m >= it.m as usize && n >= it.n as usize && k >= it.k as usize)
.max_by_key(|it| it.m as u64 * it.n as u64 * it.k as u64)
.map(|it| (it.m, it.n, it.k))
}