use burn_backend::cubecl::dtype_to_storage_type;
use burn_backend::{
DType,
ops::{ConvOptions, conv::calculate_conv_output_sizes},
};
use burn_std::{Metadata, Shape, Slice};
use core::iter;
use cubecl::{
prelude::*,
std::tensor::{TensorHandle, into_contiguous_pitched},
};
use cubek::convolution::components::ConvSetupError;
use crate::{
CubeDevice,
kernel::{
AddOp, into_contiguous_aligned, launch_binop,
matmul::{MatmulStrategy, matmul},
reduce::{KernelReduceStrategy, reduce_dim},
slice_assign, slice_with_steps,
utils::split_dim,
},
ops::{
numeric::{empty_device_dtype, zeros_client},
reshape, swap_dims,
},
tensor::CubeTensor,
};
use cubek::reduce::components::instructions::ReduceOperationConfig;
#[cfg(not(test))]
pub(crate) fn batches_per_run(
batch_size: usize,
out_shape: usize,
plane_size: usize,
) -> Result<usize, ConvSetupError> {
use cubek::matmul::definition::MatmulAvailabilityError;
let cube_count_per_batch = out_shape.div_ceil(plane_size);
let max_cube_count = u16::MAX as usize;
let max_simultaneous = Ord::min(max_cube_count / cube_count_per_batch, batch_size);
if max_simultaneous == 0 {
return Err(MatmulAvailabilityError::CubeCountTooBig(CubeCount::Static(
cube_count_per_batch as u32,
1,
1,
))
.into());
}
Ok((0..=max_simultaneous)
.rev()
.find(|per_run| batch_size.is_multiple_of(*per_run))
.expect("Logically not possible"))
}
#[cfg(test)]
#[allow(unused)]
pub(crate) fn batches_per_run(
batch_size: usize,
out_shape: usize,
plane_size: usize,
) -> Result<usize, ConvSetupError> {
Ok(1)
}
pub fn conv_im2col_1x1<const N: usize>(
input: CubeTensor,
weight: CubeTensor,
bias: Option<CubeTensor>,
options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
let rank = input.meta.num_dims();
let dim_c = rank - 1;
let out_channels = weight.meta.shape()[0];
check_pointwise_strided(&weight.meta.shape()[1..dim_c], &options)?;
let out_shape = calculate_conv_output_sizes(
&weight.meta.shape()[1..dim_c],
&options.stride,
&options.padding,
&options.dilation,
&input.meta.shape()[1..dim_c],
);
let mut split_m = vec![input.meta.shape()[0]];
split_m.extend(out_shape.iter().copied());
let input = match options.stride.iter().all(|stride| *stride == 1) {
true => input,
false => strided_spatial_view(input, &out_shape, &options.stride),
};
let input = reshape_input(input); let dtype = input.dtype;
let weight = swap_dims(reshape_weight(weight), 0, 1);
let out = matmul(input, weight, None, MatmulStrategy::default(), dtype)?;
let mut out = split_dim(out, 0, &split_m);
if let Some(bias) = bias {
let mut bias_shape = iter::repeat_n(1, rank - 1).collect::<Vec<_>>();
bias_shape.push(out_channels);
let bias = reshape(bias, bias_shape.into());
out = launch_binop::<AddOp>(out, bias);
}
Ok(out)
}
fn reshape_input(input: CubeTensor) -> CubeTensor {
let input = crate::kernel::untile(input);
let rank = input.meta.num_dims();
let dim_c = rank - 1;
let dtype = input.dtype;
let batch_size = input.meta.shape()[0];
let in_c: usize = input.meta.shape()[dim_c];
let in_shape: Shape = input.meta.shape()[1..dim_c].into();
let mut input = if !is_spatial_contiguous(input.meta.shape(), input.meta.strides()) {
let (client, device) = (input.client.clone(), input.device.clone());
let contiguous =
into_contiguous_pitched(&client, input.binding(), dtype_to_storage_type(dtype));
from_handle(client, device, contiguous, dtype)
} else {
input
};
*input.meta = Metadata::new(
[batch_size * in_shape.num_elements(), in_c], [input.meta.strides()[dim_c - 1], input.meta.strides()[dim_c]],
);
input
}
fn is_spatial_contiguous(shape: &[usize], strides: &[usize]) -> bool {
let rank = shape.len();
let dim_c = rank - 1;
if strides[dim_c] != 1 {
return false;
}
for i in (1..dim_c).rev() {
if strides[i + 1] * shape[i + 1] != strides[i] {
return false;
}
}
true
}
fn from_handle(
client: Client,
device: CubeDevice,
handle: TensorHandle,
dtype: DType,
) -> CubeTensor {
CubeTensor::new(
client.clone(),
handle.handle,
*handle.metadata,
device.clone(),
dtype,
)
}
fn check_pointwise<const N: usize>(
kernel_shape: &[usize],
options: &ConvOptions<N>,
) -> Result<(), ConvSetupError> {
check_pointwise_strided(kernel_shape, options)?;
match options.stride.iter().all(|stride| *stride == 1) {
true => Ok(()),
false => Err(ConvSetupError::Unknown),
}
}
fn check_pointwise_strided<const N: usize>(
kernel_shape: &[usize],
options: &ConvOptions<N>,
) -> Result<(), ConvSetupError> {
if options.groups != 1 {
return Err(ConvSetupError::Groups(options.groups));
}
let pointwise = kernel_shape.iter().all(|size| *size == 1)
&& options
.padding
.iter()
.all(|&(begin, end)| begin == 0 && end == 0)
&& options.dilation.iter().all(|dilation| *dilation == 1);
match pointwise {
true => Ok(()),
false => Err(ConvSetupError::Unknown),
}
}
fn strided_spatial_view(input: CubeTensor, out_shape: &[usize], stride: &[usize]) -> CubeTensor {
let mut input = crate::kernel::untile(input);
let mut shape = input.meta.shape().to_vec();
let mut strides = input.meta.strides().to_vec();
for (dim, (out, step)) in out_shape.iter().zip(stride).enumerate() {
shape[dim + 1] = *out;
strides[dim + 1] *= *step;
}
*input.meta = Metadata::new(shape, strides);
input
}
fn reshape_weight(weight: CubeTensor) -> CubeTensor {
let mut weight = crate::kernel::untile(weight);
let dim_c = weight.meta.num_dims() - 1;
let strides = [weight.meta.strides()[0], weight.meta.strides()[dim_c]];
let shape = [weight.meta.shape()[0], weight.meta.shape()[dim_c]];
*weight.meta = Metadata::new(shape, strides);
match strides[1] {
1 => weight,
_ => into_contiguous_aligned(weight),
}
}
pub fn dgrad_im2col_1x1<const N: usize>(
out_grad: CubeTensor,
weight: CubeTensor,
input_shape: Shape,
options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
let dim_c = out_grad.meta.num_dims() - 1;
check_pointwise(&weight.meta.shape()[1..dim_c], &options)?;
let split_m = input_shape[..dim_c].to_vec();
let out_grad = reshape_input(out_grad); let dtype = out_grad.dtype;
let weight = reshape_weight(weight);
let out = matmul(out_grad, weight, None, MatmulStrategy::default(), dtype)?;
Ok(split_dim(out, 0, &split_m)) }
pub fn wgrad_im2col_1x1<const N: usize>(
input: CubeTensor,
out_grad: CubeTensor,
weight_shape: Shape,
options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
let dim_c = input.meta.num_dims() - 1;
check_pointwise(&weight_shape[1..dim_c], &options)?;
let input = reshape_input(input); let out_grad = reshape_input(out_grad); let dtype = out_grad.dtype;
let out_grad = swap_dims(out_grad, 0, 1);
let grad = matmul(out_grad, input, None, MatmulStrategy::default(), dtype)?;
Ok(reshape(grad, weight_shape)) }
const MIN_SPLIT_ROWS: usize = 2048;
const MAX_SPLIT: usize = 64;
fn split_count(k: usize) -> Option<usize> {
if k < MIN_SPLIT_ROWS * 2 {
return None;
}
let by_rows = k / MIN_SPLIT_ROWS;
let ceiling = Ord::min(by_rows, MAX_SPLIT);
(1..=ceiling.ilog2())
.rev()
.map(|log| 1usize << log)
.find(|split| k.is_multiple_of(*split))
}
pub fn wgrad_im2col_1x1_split<const N: usize>(
input: CubeTensor,
out_grad: CubeTensor,
weight_shape: Shape,
options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
let dim_c = input.meta.num_dims() - 1;
check_pointwise(&weight_shape[1..dim_c], &options)?;
let uncut = {
let args = (
input.clone(),
out_grad.clone(),
weight_shape.clone(),
options.clone(),
);
move || wgrad_im2col_1x1::<N>(args.0, args.1, args.2, args.3)
};
let rows: usize = input.meta.shape()[..dim_c].iter().product();
let Some(split) = split_count(rows) else {
return uncut();
};
let per = rows / split;
let input = reshape_input(input); let out_grad = reshape_input(out_grad); let dtype = out_grad.dtype;
let in_channels = input.meta.shape()[1];
let out_channels = out_grad.meta.shape()[1];
let input = reshape(input, Shape::new([split, per, in_channels]));
let out_grad = reshape(out_grad, Shape::new([split, per, out_channels]));
let out_grad = swap_dims(out_grad, 1, 2);
let Ok(partials) = matmul(out_grad, input, None, MatmulStrategy::default(), dtype) else {
return uncut();
};
let grad = reduce_dim(
partials,
None,
0,
KernelReduceStrategy::default(),
ReduceOperationConfig::Sum,
);
let Ok(grad) = grad else {
return uncut();
};
Ok(reshape(grad, weight_shape))
}
fn im2col<const N: usize>(
input: CubeTensor,
out_shape: &[usize],
kernel_shape: &[usize],
options: &ConvOptions<N>,
) -> CubeTensor {
let rank = input.meta.num_dims();
let dim_c = rank - 1;
let batch = input.meta.shape()[0];
let channels = input.meta.shape()[dim_c];
let in_shape = input.meta.shape()[1..dim_c].to_vec();
let taps: usize = kernel_shape.iter().product();
let mut columns_shape = vec![batch];
columns_shape.extend(out_shape.iter().copied());
columns_shape.push(taps * channels);
let mut blocks = Vec::with_capacity(taps);
let mut clipped = false;
for tap in 0..taps {
let mut rest = tap;
let mut offsets = vec![0usize; N];
for axis in (0..N).rev() {
offsets[axis] = rest % kernel_shape[axis];
rest /= kernel_shape[axis];
}
let mut source = vec![Slice::from(0..batch)];
let mut target = vec![Slice::from(0..batch)];
let mut covers_nothing = false;
for axis in 0..N {
let stride = options.stride[axis] as isize;
let base = (offsets[axis] * options.dilation[axis]) as isize
- options.padding_begin()[axis] as isize;
let extent = in_shape[axis] as isize;
let first = match base >= 0 {
true => 0,
false => (-base + stride - 1) / stride,
};
let last = match extent - 1 - base {
reach if reach < 0 => 0,
reach => Ord::min(out_shape[axis] as isize, reach / stride + 1),
};
if last <= first {
covers_nothing = true;
break;
}
clipped |= first > 0 || last < out_shape[axis] as isize;
let start = first * stride + base;
source.push(Slice {
start,
end: Some(start + (last - first - 1) * stride + 1),
step: stride,
});
target.push(Slice {
start: first,
end: Some(last),
step: 1,
});
}
if covers_nothing {
clipped = true;
continue;
}
source.push(Slice::from(0..channels));
target.push(Slice::from(tap * channels..(tap + 1) * channels));
blocks.push((source, target));
}
let mut columns = match clipped {
true => zeros_client(
input.client.clone(),
input.device.clone(),
columns_shape.into(),
input.dtype,
),
false => empty_device_dtype(
input.client.clone(),
input.device.clone(),
columns_shape.into(),
input.dtype,
),
};
for (source, target) in blocks {
let block = slice_with_steps(input.clone(), &source);
columns = slice_assign(columns, &target, block);
}
columns
}
pub fn wgrad_im2col<const N: usize>(
input: CubeTensor,
out_grad: CubeTensor,
weight_shape: Shape,
options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
let rank = input.meta.num_dims();
let dim_c = rank - 1;
if options.groups != 1 {
return Err(ConvSetupError::Groups(options.groups));
}
if check_pointwise(&weight_shape[1..dim_c], &options).is_ok() {
return Err(ConvSetupError::Unknown);
}
let out_channels = weight_shape[0];
let in_channels = input.meta.shape()[dim_c];
let kernel_shape = weight_shape[1..dim_c].to_vec();
let out_shape = out_grad.meta.shape()[1..dim_c].to_vec();
let cols = kernel_shape.iter().product::<usize>() * in_channels;
let rows = out_grad.meta.shape()[..dim_c].iter().product::<usize>();
let columns = im2col::<N>(input, &out_shape, &kernel_shape, &options);
let columns = reshape(columns, Shape::new([rows, cols]));
let out_grad = reshape_input(out_grad); let dtype = out_grad.dtype;
let uncut = |columns: CubeTensor, out_grad: CubeTensor| {
let out_grad = swap_dims(out_grad, 0, 1); matmul(out_grad, columns, None, MatmulStrategy::default(), dtype)
};
let grad = match split_count(rows) {
Some(split) => {
let per = rows / split;
let cut = reshape(columns.clone(), Shape::new([split, per, cols]));
let grad = reshape(out_grad.clone(), Shape::new([split, per, out_channels]));
let grad = swap_dims(grad, 1, 2);
let partials = matmul(grad, cut, None, MatmulStrategy::default(), dtype)
.ok()
.and_then(|partials| {
reduce_dim(
partials,
None,
0,
KernelReduceStrategy::default(),
ReduceOperationConfig::Sum,
)
.ok()
});
match partials {
Some(grad) => grad,
None => uncut(columns, out_grad)?,
}
}
None => uncut(columns, out_grad)?,
};
Ok(reshape(grad, weight_shape))
}