use super::init_matmul_output;
use crate::{kernel::quantization::dequantize, tensor::CubeTensor};
use burn_backend::cubecl::dtype_to_storage_type;
use burn_backend::{DType, TensorMetadata};
use burn_std::{MatmulTransformAnalysis, MatmulTransformPolicy};
use cubek::{
matmul::{
definition::{MatmulElems, MatmulGlobalElems, MatmulSetupError},
strategy::Strategy,
},
std::InputBinding,
};
#[cfg(feature = "autotune")]
use super::matmul_autotune;
pub enum MatmulStrategy {
#[cfg(feature = "autotune")]
Autotune,
Cube,
}
impl Default for MatmulStrategy {
fn default() -> Self {
#[cfg(feature = "autotune")]
return MatmulStrategy::Autotune;
#[cfg(not(feature = "autotune"))]
MatmulStrategy::Cube
}
}
fn is_two_level(tensor: &CubeTensor) -> bool {
match tensor.dtype {
DType::QFloat(scheme) => burn_backend::quantization::global_scale_dtype(&scheme).is_some(),
_ => false,
}
}
fn maybe_dequantize(tensor: CubeTensor, dtype: DType) -> CubeTensor {
if is_two_level(&tensor) {
dequantize(tensor, dtype)
} else {
tensor
}
}
pub fn matmul(
lhs: CubeTensor,
rhs: CubeTensor,
out: Option<CubeTensor>,
strategy: MatmulStrategy,
out_dtype: DType,
) -> Result<CubeTensor, MatmulSetupError> {
let out = out.unwrap_or_else(|| init_matmul_output(&lhs, &rhs, out_dtype));
let mut lhs = maybe_dequantize(lhs, out_dtype);
let rhs = maybe_dequantize(rhs, out_dtype);
let mut out_launch = out.clone();
if lhs.qparams.is_none() {
let analysis = MatmulTransformAnalysis::from_metadata(&lhs.meta, &rhs.meta, &out.meta);
let action = MatmulTransformPolicy::default().action(&analysis);
action.apply(&mut lhs.meta);
action.apply(&mut out_launch.meta);
}
match strategy {
MatmulStrategy::Cube => {
launch_matmul(&Strategy::default(), lhs, rhs, out_launch)?;
Ok(out)
}
#[cfg(feature = "autotune")]
MatmulStrategy::Autotune => {
matmul_autotune(lhs, rhs, Some(out_launch), out_dtype);
Ok(out)
}
}
}
pub(crate) fn launch_matmul_naive<S: Clone + Into<Strategy>>(
strategy: &S,
mut lhs: CubeTensor,
mut rhs: CubeTensor,
out: CubeTensor,
) -> Result<(), MatmulSetupError> {
if lhs.qparams.is_some() || rhs.qparams.is_some() {
match launch_matmul(strategy, lhs.clone(), rhs.clone(), out.clone()) {
Err(_) => {
if lhs.qparams.is_some() {
lhs = dequantize(lhs, out.dtype);
}
if rhs.qparams.is_some() {
rhs = dequantize(rhs, out.dtype);
}
launch_matmul(strategy, lhs, rhs, out)
}
Ok(_) => Ok(()),
}
} else {
launch_matmul(strategy, lhs, rhs, out)
}
}
pub(crate) fn launch_matmul<S: Clone + Into<Strategy>>(
strategy: &S,
lhs: CubeTensor,
mut rhs: CubeTensor,
out: CubeTensor,
) -> Result<(), MatmulSetupError> {
let strategy: Strategy = strategy.clone().into();
let client = &out.client;
let lhs_quant_handles = lhs.quantized_handles();
let out_dtype: DType = out.dtype;
let (lhs_dtype, lhs_handle) = match lhs_quant_handles {
None => {
let lhs_dtype = lhs.dtype;
(
lhs_dtype,
InputBinding::new(lhs.binding(), dtype_to_storage_type(lhs_dtype)),
)
}
Some((data, scale)) => {
let scheme = lhs.scheme();
let data_dtype = data.dtype;
let scale_dtype = scale.dtype;
(
out_dtype,
InputBinding::quantized(
data.binding(),
scale.binding(),
lhs.meta.shape().clone(),
scheme,
dtype_to_storage_type(data_dtype),
dtype_to_storage_type(scale_dtype),
),
)
}
};
let rhs_quant_handles = rhs.quantized_handles();
let (rhs_dtype, rhs_handle) = match rhs_quant_handles {
None => (
lhs_dtype,
InputBinding::new(rhs.binding(), dtype_to_storage_type(lhs_dtype)),
),
Some((data, scale)) => {
if matches!(
strategy,
Strategy::MultiLevel(cubek::matmul::multi_level::Strategy::Naive)
) && rhs.scheme().block_size().is_some()
{
rhs = dequantize(rhs.clone(), lhs_dtype);
let rhs_dtype = rhs.dtype;
(
lhs_dtype,
InputBinding::new(rhs.binding(), dtype_to_storage_type(rhs_dtype)),
)
} else {
let scheme = rhs.scheme();
let data_dtype = data.dtype;
let scale_dtype = scale.dtype;
(
out_dtype,
InputBinding::quantized(
data.binding(),
scale.binding(),
rhs.meta.shape().clone(),
scheme,
dtype_to_storage_type(data_dtype),
dtype_to_storage_type(scale_dtype),
),
)
}
}
};
let mut dtypes = MatmulElems::from_globals(&MatmulGlobalElems {
lhs: dtype_to_storage_type(lhs_dtype),
rhs: dtype_to_storage_type(rhs_dtype),
out: dtype_to_storage_type(out_dtype),
});
cubek::matmul::launch::launch_ref(
&strategy,
client,
lhs_handle,
rhs_handle,
out.clone().binding(),
&mut dtypes,
)?;
Ok(())
}