use crate::{ops::numeric::empty_device_dtype, tensor::CubeTensor};
use burn_backend::cubecl::dtype_to_storage_type;
use burn_backend::ops::{ConvOptions, conv::calculate_conv_output_sizes};
use cubek::convolution::{
ConvolutionArgs, DepthwiseStrategy, DepthwiseTensors, components::ConvSetupError,
launch_depthwise,
};
fn skip_large_filter(in_channels: usize, filter_shape: &[usize]) -> bool {
in_channels < 32 && filter_shape.iter().product::<usize>() > 256
}
pub fn conv_depthwise<const N: usize>(
input: CubeTensor,
weight: CubeTensor,
bias: Option<CubeTensor>,
options: ConvOptions<N>,
strategy: DepthwiseStrategy,
) -> Result<CubeTensor, ConvSetupError> {
if N != 2 {
return Err(ConvSetupError::Unknown);
}
if bias.is_some() {
return Err(ConvSetupError::Unknown);
}
let out_dtype = input.dtype;
let rank = input.meta.shape().num_dims();
let batch_size = input.meta.shape()[0];
let dim_c = rank - 1;
let shape = &input.meta.shape()[1..dim_c];
let out_channels = weight.meta.shape()[0];
let weight_shape = &weight.meta.shape()[1..dim_c];
if skip_large_filter(input.meta.shape()[dim_c], weight_shape) {
return Err(ConvSetupError::Unknown);
}
let mut out_shape = calculate_conv_output_sizes(
weight_shape,
&options.stride,
&options.padding,
&options.dilation,
shape,
);
out_shape.insert(0, batch_size);
out_shape.push(out_channels);
let out = empty_device_dtype(
input.client.clone(),
input.device.clone(),
out_shape.into(),
out_dtype,
);
let padding = options.padding_begin();
let args = ConvolutionArgs::<2> {
stride: [options.stride[0], options.stride[1]],
padding: [padding[0], padding[1]],
dilation: [options.dilation[0], options.dilation[1]],
};
let client = input.client.clone();
let dtype = dtype_to_storage_type(out_dtype);
let tensors = DepthwiseTensors {
input: input.binding(),
weight: weight.binding(),
out: out.clone().binding(),
};
launch_depthwise(&client, tensors, args, options.groups, dtype, strategy)?;
Ok(out)
}