use burn_backend::cubecl::{autotune::with_roofline_bounds, dtype_to_storage_type};
use cubecl::{std::throughput::roofline_bounds, tune::TunableSet};
use cubek::matmul::{
definition::{MatmulCost, MatmulGlobalElems},
tune_key::MatmulAutotuneKey,
};
use crate::kernel::matmul::tune::base::Inputs;
type MatmulTunables<Out> = TunableSet<MatmulAutotuneKey, Inputs, Out>;
pub(super) fn with_matmul_bounds<Out: 'static>(set: MatmulTunables<Out>) -> MatmulTunables<Out> {
with_roofline_bounds(set, |_key, tensors: &Inputs, thresholds| {
let client = &tensors.0.client;
let cost = cost(tensors);
roofline_bounds(client, cost.compute_key(client), cost.work(), thresholds)
})
}
fn cost((lhs, rhs, out): &Inputs) -> MatmulCost {
let lhs_shape = lhs.meta.shape();
let rhs_shape = rhs.meta.shape();
let ndims = lhs_shape.len();
MatmulCost {
batches: lhs_shape[..ndims - 2].iter().product(),
m: lhs_shape[ndims - 2],
k: lhs_shape[ndims - 1],
n: rhs_shape[ndims - 1],
elems: MatmulGlobalElems {
lhs: dtype_to_storage_type(lhs.dtype),
rhs: dtype_to_storage_type(rhs.dtype),
out: dtype_to_storage_type(out.dtype),
},
}
}