use cubecl::{
calculate_cube_count_elemwise,
client::Client,
ir::{FloatKind, VectorRegisters},
num_traits::Zero,
prelude::*,
std::tensor::layout::linear::{LinearViewMut, linear_view},
std::{FastDivmod, FastDivmodInt},
tensor_vector_size_parallel,
};
use crate::{components::ConvSetupError, launch::ConvolutionArgs};
#[cube]
fn decompose_linear<I: FastDivmodInt>(pos: I, shape: &Sequence<FastDivmod<I>>) -> (I, Sequence<I>) {
let rank = comptime![shape.len()];
let mut offs = pos;
let mut out = Sequence::new();
#[unroll]
for i in 0..rank {
let dim = comptime![rank - i - 1];
let (rem, offs_local) = shape.index(dim).div_mod(offs);
out.push(offs_local);
offs = rem;
}
(offs, out.reversed())
}
#[derive(CubeLaunch, CubeType, Clone)]
pub(crate) struct ConvParam {
pub stride: u32,
pub dilation: u32,
pub padding: i32,
}
#[derive(CubeLaunch, CubeType)]
struct Conv2dArgs {
conv_params: Sequence<ConvParam>,
channels_per_group: u32,
}
#[cube(launch_unchecked, address_type = "dynamic")]
#[allow(clippy::redundant_closure)]
fn direct_conv2d_kernel<E: Numeric, NIn: Size, NOut: Size>(
input: &Tensor<Vector<E, NIn>>,
weight: &Tensor<Vector<E, NIn>>,
bias: ComptimeOption<&[Vector<E, NOut>]>,
mut output: LinearViewMut<'_, Vector<E, NOut>>,
args: Conv2dArgs,
shape_out: Sequence<FastDivmod<u32>>,
shape_out_c: FastDivmod<u32>,
#[comptime] has_padding: bool,
#[comptime] accumulate_components: bool,
#[comptime] channel_block: usize,
#[define(E)] _dtype: ElemType,
) {
if !output.is_in_bounds(ABSOLUTE_POS) {
terminate!();
}
let n_spatial = comptime![shape_out.len()];
let vector_size_out = output.vector_size();
let pos = ABSOLUTE_POS * vector_size_out;
let in_c_per_group = weight.shape(weight.rank() - 1) as u32;
let (rem, out_c) = shape_out_c.div_mod(pos as u32);
let (b, spatial_pos) = decompose_linear(rem, &shape_out);
let g = out_c / args.channels_per_group;
let ic_start = in_c_per_group * g;
let bias: ComptimeOption<Vector<E, NOut>> =
bias.map(|bias| bias[out_c as usize / vector_size_out]);
let mut sum = bias.unwrap_or_else(|| Vector::zero());
let in_offs = b as usize * input.stride(0) + ic_start as usize;
let stride_oc = weight.stride(0);
let mut in_shape = Sequence::new();
let mut in_strides = Sequence::new();
let mut kernel_shape = Sequence::new();
let mut kernel_strides = Sequence::new();
#[unroll]
for i in 0..n_spatial {
in_shape.push(input.shape(i + 1) as u32);
in_strides.push(input.stride(i + 1));
kernel_shape.push(weight.shape(i + 1) as u32);
kernel_strides.push(weight.stride(i + 1));
}
let weight_offs = out_c as usize * stride_oc;
let loop_params = LoopParams {
out_pos: spatial_pos,
in_shape,
in_strides,
kernel_shape,
kernel_strides,
conv_params: args.conv_params,
in_c_per_group,
stride_oc,
};
let vector_size_in = input.vector_size();
if accumulate_components {
comptime!(assert!(
vector_size_out.is_multiple_of(channel_block),
"a block of {channel_block} channels does not divide the output vector of {vector_size_out}"
));
#[unroll]
for bi in 0..comptime![vector_size_out / channel_block] {
let base_v = bi * channel_block;
let mut partials = Array::<Vector<E, NIn>>::new(channel_block);
kernel_loop(
input,
weight,
&mut sum,
&mut partials,
in_offs,
true,
weight_offs,
&loop_params,
0usize,
has_padding,
accumulate_components,
base_v,
);
#[unroll]
for j in 0..channel_block {
let mut channel = sum.extract(base_v + j);
#[unroll]
for i in 0..vector_size_in {
channel += partials[j].extract(i);
}
sum.insert(base_v + j, channel);
}
}
} else {
let mut partials = Array::<Vector<E, NIn>>::new(1usize);
kernel_loop(
input,
weight,
&mut sum,
&mut partials,
in_offs,
true,
weight_offs,
&loop_params,
0usize,
has_padding,
accumulate_components,
0usize,
);
}
output.write(ABSOLUTE_POS, sum);
}
#[derive(CubeType, Clone)]
struct LoopParams {
out_pos: Sequence<u32>,
in_shape: Sequence<u32>,
in_strides: Sequence<usize>,
kernel_shape: Sequence<u32>,
kernel_strides: Sequence<usize>,
conv_params: Sequence<ConvParam>,
in_c_per_group: u32,
stride_oc: usize,
}
#[cube]
fn kernel_loop<E: Numeric, NIn: Size, NOut: Size>(
input: &Tensor<Vector<E, NIn>>,
weight: &Tensor<Vector<E, NIn>>,
sum: &mut Vector<E, NOut>,
partials: &mut Array<Vector<E, NIn>>,
in_offs: usize,
in_bounds: bool,
weight_offs: usize,
params: &LoopParams,
#[comptime] kernel_dim: usize,
#[comptime] has_padding: bool,
#[comptime] accumulate_components: bool,
#[comptime] base_v: usize,
) {
if comptime![kernel_dim < params.kernel_shape.len()] {
let out_idx = *params.out_pos.index(kernel_dim);
let conv = params.conv_params.index(kernel_dim);
let shape = *params.in_shape.index(kernel_dim);
let stride = *params.in_strides.index(kernel_dim);
let k_stride = *params.kernel_strides.index(kernel_dim);
for pos in 0..*params.kernel_shape.index(kernel_dim) {
let in_pos = (out_idx * conv.stride + pos * conv.dilation) as i32 - conv.padding;
let in_offs = in_offs + in_pos as usize * stride;
let weight_offs = weight_offs + pos as usize * k_stride;
let mut in_bounds = in_bounds;
if has_padding {
in_bounds &= in_pos >= 0 && (in_pos as u32) < shape;
}
kernel_loop(
input,
weight,
sum,
partials,
in_offs,
in_bounds,
weight_offs,
params,
comptime![kernel_dim + 1],
has_padding,
accumulate_components,
base_v,
);
}
} else {
kernel_loop_inner(
input,
weight,
sum,
partials,
in_offs,
in_bounds,
weight_offs,
params.in_c_per_group,
params.stride_oc,
accumulate_components,
base_v,
);
}
}
#[cube]
fn kernel_loop_inner<E: Numeric, NIn: Size, NOut: Size>(
input: &Tensor<Vector<E, NIn>>,
weight: &Tensor<Vector<E, NIn>>,
sum: &mut Vector<E, NOut>,
partials: &mut Array<Vector<E, NIn>>,
in_offs: usize,
in_bounds: bool,
weight_offs: usize,
in_c_per_group: u32,
stride_oc: usize,
#[comptime] accumulate_components: bool,
#[comptime] base_v: usize,
) {
if in_bounds {
if accumulate_components {
accumulate_in_components(
input,
weight,
partials,
in_offs,
weight_offs,
in_c_per_group,
stride_oc,
base_v,
);
} else {
accumulate_per_step(
input,
weight,
sum,
in_offs,
weight_offs,
in_c_per_group,
stride_oc,
);
}
}
}
#[cube]
fn accumulate_in_components<E: Numeric, NIn: Size>(
input: &Tensor<Vector<E, NIn>>,
weight: &Tensor<Vector<E, NIn>>,
partials: &mut Array<Vector<E, NIn>>,
in_offs: usize,
weight_offs: usize,
in_c_per_group: u32,
stride_oc: usize,
#[comptime] base_v: usize,
) {
let vector_size_in = input.vector_size();
let block = partials.len();
for in_c in range_stepped(0, in_c_per_group, vector_size_in as u32) {
let val = input[(in_offs + in_c as usize) / vector_size_in];
#[unroll]
for j in 0..block {
let weight_offs = weight_offs + (base_v + j) * stride_oc + in_c as usize;
partials[j] += val * weight[weight_offs / vector_size_in];
}
}
}
#[cube]
fn accumulate_per_step<E: Numeric, NIn: Size, NOut: Size>(
input: &Tensor<Vector<E, NIn>>,
weight: &Tensor<Vector<E, NIn>>,
sum: &mut Vector<E, NOut>,
in_offs: usize,
weight_offs: usize,
in_c_per_group: u32,
stride_oc: usize,
) {
let vector_size_in = input.vector_size();
let vector_size_out = sum.vector_size();
for in_c in range_stepped(0, in_c_per_group, vector_size_in as u32) {
let in_pos = in_offs + in_c as usize;
let mut weight_pos = weight_offs + in_c as usize;
let val = input[in_pos / vector_size_in];
#[unroll]
for v in 0..vector_size_out {
let weight = weight[weight_pos / vector_size_in];
let val = val * weight;
#[unroll]
for i in 0..vector_size_in {
sum.insert(v, sum.extract(v) + val.extract(i));
}
weight_pos += stride_oc;
}
}
}
pub struct DirectTensors {
pub input: TensorBinding,
pub weight: TensorBinding,
pub bias: Option<TensorBinding>,
pub out: TensorBinding,
}
pub fn launch_direct<const N: usize>(
client: &Client,
tensors: DirectTensors,
args: ConvolutionArgs<N>,
groups: usize,
dtype: ElemType,
) -> Result<(), ConvSetupError> {
let DirectTensors {
input,
weight,
bias,
out,
} = tensors;
let rank = input.shape.len();
let dim_c = rank - 1;
let in_shape = &input.shape[1..dim_c];
let out_channels = weight.shape[0];
let kernel_shape = &weight.shape[1..dim_c];
let out_size = &out.shape[1..dim_c];
let channels_per_group = out_channels / groups;
let check_spatial_bounds = should_check_spatial_bounds(in_shape, kernel_shape, out_size, &args);
let mut grouped_out_shape = out.shape.clone();
grouped_out_shape[dim_c] = channels_per_group;
let vector_size_out = tensor_vector_size_parallel(
client.io_optimized_vector_sizes(dtype.size()),
&grouped_out_shape,
&out.strides,
dim_c,
);
let vector_size_in = tensor_vector_size_parallel(
client.io_optimized_vector_sizes(dtype.size()),
&weight.shape,
&weight.strides,
weight.shape.len() - 1,
);
let accumulate_components =
client.properties().hardware.plane_size_max == 1 && vector_size_in > 1;
let channel_block =
VectorRegisters::new(&client.properties().hardware, register_elem_size(dtype))
.map_or(1, |registers| {
channel_block(registers, vector_size_in, vector_size_out)
});
let shape_out = out.shape[1..dim_c].iter().map(|s| *s as u32).collect();
let shape_out_c = out_channels as u32;
let mut conv_params = SequenceArg::new();
for i in 0..kernel_shape.len() {
conv_params.push(ConvParamLaunch::new(
args.stride[i] as u32,
args.dilation[i] as u32,
args.padding[i] as i32,
));
}
let working_units = out.shape.iter().product::<usize>() / vector_size_out as usize;
let cube_dim = CubeDim::new(client, working_units);
let cube_count = calculate_cube_count_elemwise(client, working_units, cube_dim);
let address_type = input
.required_address_type(dtype.size())
.max(weight.required_address_type(dtype.size()))
.max(out.required_address_type(dtype.size()));
unsafe {
direct_conv2d_kernel::launch_unchecked(
client,
cube_count,
cube_dim,
address_type,
vector_size_in,
vector_size_out,
input.into_tensor_arg(),
weight.into_tensor_arg(),
bias.map(|b| b.into_buffer_arg()).into(),
linear_view(out),
Conv2dArgsLaunch::new(conv_params, channels_per_group as u32),
shape_out,
shape_out_c,
check_spatial_bounds,
accumulate_components,
channel_block,
dtype,
)
};
Ok(())
}
fn channel_block(
registers: VectorRegisters,
vector_size_in: usize,
vector_size_out: usize,
) -> usize {
let operands = registers.count() / 2;
let block = registers.vectors_fitting(vector_size_in, operands).max(1);
(1 << block.ilog2()).min(vector_size_out)
}
fn register_elem_size(dtype: ElemType) -> usize {
match dtype {
ElemType::Float(FloatKind::F16 | FloatKind::BF16) => size_of::<f32>(),
dtype => dtype.size(),
}
}
fn should_check_spatial_bounds<const N: usize>(
in_shape: &[usize],
kernel_shape: &[usize],
out_shape: &[usize],
args: &ConvolutionArgs<N>,
) -> bool {
(0..N).any(|dim| {
let begin = args.padding[dim] as i64;
let first = -begin;
let last = (out_shape[dim] as i64 - 1) * args.stride[dim] as i64
+ (kernel_shape[dim] as i64 - 1) * args.dilation[dim] as i64
- begin;
first < 0 || last >= in_shape[dim] as i64
})
}
#[cfg(test)]
mod tests {
use cubek_test_utils::hardware::{AVX2, AVX512, NEON};
use super::*;
fn avx2(elem_size: usize) -> VectorRegisters {
VectorRegisters::new(&AVX2, elem_size).unwrap()
}
fn avx512(elem_size: usize) -> VectorRegisters {
VectorRegisters::new(&AVX512, elem_size).unwrap()
}
fn neon(elem_size: usize) -> VectorRegisters {
VectorRegisters::new(&NEON, elem_size).unwrap()
}
#[test]
fn a_two_register_accumulator_halves_the_block() {
assert_eq!(channel_block(avx2(4), 16, 16), 4);
}
#[test]
fn a_four_register_accumulator_quarters_the_block() {
assert_eq!(channel_block(neon(4), 16, 16), 4);
}
#[test]
fn a_register_wide_accumulator_spends_one() {
assert_eq!(channel_block(avx2(4), 8, 16), 8);
assert_eq!(channel_block(avx512(4), 16, 16), 16);
}
#[test]
fn a_narrow_accumulator_still_spends_a_whole_register() {
assert_eq!(channel_block(avx2(4), 2, 16), 8);
}
#[test]
fn a_half_float_accumulator_spends_f32_registers() {
let f16 = register_elem_size(ElemType::Float(FloatKind::F16));
assert_eq!(channel_block(avx2(f16), 32, 32), 2);
assert_eq!(channel_block(neon(f16), 32, 32), 2);
}
#[test]
fn the_block_never_outgrows_the_output_vector() {
assert_eq!(channel_block(avx2(4), 8, 2), 2);
}
#[test]
fn an_accumulator_wider_than_the_budget_still_gets_a_block_of_one() {
assert_eq!(channel_block(avx2(8), 64, 16), 1);
}
}