burn-cubecl 0.22.0

Generic backend that can be compiled just-in-time to any shader language target
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>;

/// Registers the performance bounds used for matrix multiplication autotuning.
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),
        },
    }
}