use cubecl::{
prelude::*,
server::LaunchError,
zspace::{Shape, Strides},
};
use cubek_std::InputBinding;
use cubek_tile::kind::Boundary;
use cubek_tile::launch::Grid;
use cubek_tile::layout::PhysicalAxisMap;
use cubek_tile::*;
use crate::{components::ConvSetupError, launch::ConvolutionArgs};
const REGISTER_BLOCK: RegisterBlock = RegisterBlock::new(64).split_edge();
const B: Axis = Axis(0);
const OH: Axis = Axis(1);
const OW: Axis = Axis(2);
const C: Axis = Axis(3);
const RH: Axis = Axis(4);
const RW: Axis = Axis(5);
#[cfg(test)]
const LABELS: [(Axis, &str); 6] = [
(B, "b"),
(OH, "oh"),
(OW, "ow"),
(C, "c"),
(RH, "rh"),
(RW, "rw"),
];
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct DepthwiseSpace {
b: usize,
oh: usize,
ow: usize,
c: usize,
rh: usize,
rw: usize,
rows: usize,
cols: usize,
tile_c: usize,
width: usize,
plane_size: usize,
}
impl DepthwiseSpace {
pub fn extents(&self) -> Vec<(Axis, usize)> {
vec![
(B, self.b),
(OH, self.oh),
(OW, self.ow),
(C, self.c),
(RH, self.rh),
(RW, self.rw),
]
}
pub fn levels(&self) -> Vec<Level> {
let Self {
rows,
cols,
tile_c,
width,
plane_size,
..
} = *self;
let plane_c = width * plane_size;
assert!(
tile_c.is_multiple_of(plane_c),
"DepthwiseSpace: {plane_size} units of {width} channels do not divide a tile of {tile_c}"
);
let plane_units = Levels::leaf(&[(C, width), (OW, cols), (OH, 1)])
.units(&[(C, plane_size)])
.interleaved(C);
let lines = match tile_c / plane_c {
1 => plane_units,
further => plane_units.walk(&[(C, further)]),
};
lines
.planes(&[(OH, rows)])
.cubes(&[C, OW, OH])
.batches(&[B])
.build()
}
pub fn space(&self) -> Space {
Space::new(&self.extents())
}
pub fn partitioning(&self) -> Partitioning {
Partitioning::new(self.space(), self.levels())
}
pub fn grid(&self) -> (CubeCount, CubeDim) {
(
CubeCount::Static(
self.c.div_ceil(self.tile_c) as u32,
self.ow.div_ceil(self.cols) as u32,
(self.oh.div_ceil(self.rows) * self.b) as u32,
),
CubeDim::new_2d(self.plane_size as u32, self.rows as u32),
)
}
}
#[cube(launch)]
fn depthwise_kernel<E: Numeric, V: Size>(
weight: &TileArg<'_, E, V>,
input: &TileArg<'_, E, V>,
out: &TileArg<'_, E, V>,
space: Partitioning,
#[define(E)] _dtype: ElemType,
) {
let weight = weight.tile(comptime!(space.clone()));
let input = input.tile(comptime!(space.clone()));
let out = out.tile(comptime!(space.clone()));
for cube in space {
let out = out.at(&cube);
let weight = weight.at(&cube);
let input = input.at(&cube);
for plane in cube {
for unit in plane.leaves() {
let mut out = out
.at(&unit)
.accumulating(REGISTER_BLOCK, Semiring::SUM_PROD);
out.mm(&weight.at(&unit), &input.at(&unit));
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct DepthwiseTiling {
pub rows: usize,
pub cols: usize,
pub chans: usize,
pub lines: usize,
}
impl Default for DepthwiseTiling {
fn default() -> Self {
Self {
rows: 4,
cols: 4,
chans: 1,
lines: 1,
}
}
}
impl DepthwiseTiling {
const INSTRUCTION_BOUND_TAPS: usize = 25;
const WIDE_BLOCK_UNIT_MULTIPLE: usize = 8;
pub fn for_problem(channels: usize, taps: usize, plane_units: usize) -> Self {
let deep_window = taps >= Self::INSTRUCTION_BOUND_TAPS;
let wide_block = channels >= Self::WIDE_BLOCK_UNIT_MULTIPLE * plane_units;
Self {
lines: match deep_window && wide_block {
true => 4,
false => 1,
},
..Default::default()
}
}
fn validate(self) -> Result<Self, ConvSetupError> {
if self.rows == 0 || self.cols == 0 || self.chans == 0 || self.lines == 0 {
return Err(ConvSetupError::InvalidConfig(Box::new(format!(
"depthwise tiling dimensions must be non-zero, got rows {}, cols {}, chans {}, \
lines {}",
self.rows, self.cols, self.chans, self.lines
))));
}
Ok(self)
}
fn channel_tile(self, plane_units: usize, width: usize) -> Result<usize, ConvSetupError> {
let tile = plane_units
.checked_mul(width)
.and_then(|tile| tile.checked_mul(self.chans))
.ok_or_else(|| {
ConvSetupError::InvalidConfig(Box::new(format!(
"depthwise channel tile overflows: {plane_units} units * {width} channels/line * {} \
lines/unit",
self.chans
)))
})?;
if tile == 0 {
return Err(ConvSetupError::InvalidConfig(Box::new(format!(
"depthwise channel tile must be non-zero, got {plane_units} units * {width} \
channels/line * {} lines/unit",
self.chans
))));
}
Ok(tile)
}
fn plan(
&self,
geometry: &Geometry,
plane_units: usize,
tile_c: usize,
width: usize,
) -> DepthwiseSpace {
DepthwiseSpace {
b: geometry.b,
oh: geometry.oh,
ow: geometry.ow,
c: geometry.c,
rh: geometry.rh,
rw: geometry.rw,
rows: self.rows,
cols: self.cols,
tile_c,
width,
plane_size: plane_units,
}
}
}
pub struct DepthwiseTensors {
pub input: TensorBinding,
pub weight: TensorBinding,
pub out: TensorBinding,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum DepthwiseStrategy {
Routine,
Fixed(DepthwiseTiling),
}
pub fn launch_depthwise(
client: &Client,
tensors: DepthwiseTensors,
args: ConvolutionArgs<2>,
groups: usize,
dtype: ElemType,
strategy: DepthwiseStrategy,
) -> Result<(), ConvSetupError> {
let geometry = Geometry::new(&tensors, args, groups)?;
let plane_units = plane_units(client);
let tiling = match strategy {
DepthwiseStrategy::Routine => {
DepthwiseTiling::for_problem(geometry.c, geometry.taps(), plane_units)
}
DepthwiseStrategy::Fixed(tiling) => tiling,
}
.validate()?;
let DepthwiseTensors { input, weight, out } = tensors;
let weight = geometry
.channels_innermost(client, weight, dtype)
.map_err(|_| ConvSetupError::Unknown)?;
let width = line_width(
client,
geometry.c,
dtype,
tiling.lines,
&[&input, &weight, &out],
);
let tile_c = tiling.channel_tile(plane_units, width)?;
let plan = tiling.plan(&geometry, plane_units, tile_c, width);
let launch = {
let partitioning = plan.partitioning();
let concrete = partitioning.space().clone();
let (cube_count, cube_dim) = plan.grid();
Launcher::new(
client,
partitioning,
&concrete,
Grid::Stated {
cube_count,
cube_dim,
},
)?
};
let ragged_c = !geometry.c.is_multiple_of(tile_c);
let ragged_oh = !geometry.oh.is_multiple_of(tiling.rows);
let ragged_ow = !geometry.ow.is_multiple_of(tiling.cols);
let check_h = geometry.should_check_height_bounds();
let check_w = geometry.should_check_width_bounds();
let [ph, pw] = geometry.padding;
let [sh, sw] = geometry.stride;
let [dh, dw] = geometry.dilation;
let in_spec = TileSpec::new(Projection::new(
&[B, OH, OW, RH, RW, C],
&[
PhysicalAxisMap::of(B),
PhysicalAxisMap::affine(&[(OH, sh), (RH, dh)]).shifted(-(ph as isize)),
PhysicalAxisMap::affine(&[(OW, sw), (RW, dw)]).shifted(-(pw as isize)),
PhysicalAxisMap::of(C),
],
))
.boundaries(&[
None,
guard(check_h || ragged_oh),
guard(check_w || ragged_ow),
guard(ragged_c),
]);
let w_spec = TileSpec::direct(&[RH, RW, C]).boundaries(&[None, None, guard(ragged_c)]);
let out_spec = TileSpec::direct(&[B, OH, OW, C]).boundaries(&[
None,
guard(ragged_oh),
guard(ragged_ow),
guard(ragged_c),
]);
depthwise_kernel::launch(
client,
launch.cube_count(),
launch.cube_dim(),
width,
TileArgLaunch::new(weight.into_tensor_arg(), w_spec),
TileArgLaunch::new(input.into_tensor_arg(), in_spec),
TileArgLaunch::new(out.into_tensor_arg(), out_spec),
launch.partitioning_arg(),
dtype,
);
Ok(())
}
struct Geometry {
b: usize,
ih: usize,
iw: usize,
oh: usize,
ow: usize,
c: usize,
rh: usize,
rw: usize,
stride: [usize; 2],
padding: [usize; 2],
dilation: [usize; 2],
}
impl Geometry {
fn new(
tensors: &DepthwiseTensors,
args: ConvolutionArgs<2>,
groups: usize,
) -> Result<Self, ConvSetupError> {
let input_channels = tensors.input.shape[3];
let output_channels = tensors.out.shape[3];
let weight_channels = tensors.weight.shape[0];
let weight_group_channels = tensors.weight.shape[3];
if groups != input_channels
|| output_channels != input_channels
|| weight_channels != input_channels
|| weight_group_channels != 1
{
return Err(ConvSetupError::NotDepthwise {
groups,
input_channels,
output_channels,
weight_channels,
weight_group_channels,
});
}
Ok(Self {
b: tensors.out.shape[0],
ih: tensors.input.shape[1],
iw: tensors.input.shape[2],
oh: tensors.out.shape[1],
ow: tensors.out.shape[2],
c: input_channels,
rh: tensors.weight.shape[1],
rw: tensors.weight.shape[2],
stride: args.stride,
padding: args.padding,
dilation: args.dilation,
})
}
fn taps(&self) -> usize {
self.rh * self.rw
}
fn should_check_height_bounds(&self) -> bool {
spatial_bounds_required(
self.ih,
self.oh,
self.rh,
self.stride[0],
self.padding[0],
self.dilation[0],
)
}
fn should_check_width_bounds(&self) -> bool {
spatial_bounds_required(
self.iw,
self.ow,
self.rw,
self.stride[1],
self.padding[1],
self.dilation[1],
)
}
fn channels_innermost(
&self,
client: &Client,
weight: TensorBinding,
dtype: ElemType,
) -> Result<TensorBinding, LaunchError> {
let mut permuted = weight;
let channel_stride = permuted.strides[0];
let row_stride = permuted.strides[1];
let col_stride = permuted.strides[2];
permuted.shape = Shape::from(vec![self.rh, self.rw, self.c]);
permuted.strides = Strides::new(&[row_stride, col_stride, channel_stride]);
Ok(InputBinding::new(permuted, dtype)
.into_contiguous(client)?
.into_data())
}
}
fn spatial_bounds_required(
input_size: usize,
output_size: usize,
kernel_size: usize,
stride: usize,
padding: usize,
dilation: usize,
) -> bool {
let first = -(padding as i64);
let last = (output_size as i64 - 1) * stride as i64
+ (kernel_size as i64 - 1) * dilation as i64
- padding as i64;
first < 0 || last >= input_size as i64
}
fn plane_units(client: &Client) -> usize {
client.properties().hardware.plane_size_max as usize
}
fn guard(ragged: bool) -> Option<Boundary> {
ragged.then_some(Boundary::Zero)
}
fn line_width(
client: &Client,
channels: usize,
dtype: ElemType,
requested: usize,
operands: &[&TensorBinding],
) -> usize {
if !operands.iter().all(|b| b.strides.last() == Some(&1)) {
return 1;
}
client
.io_optimized_vector_sizes(dtype.size())
.filter(|&v| {
v <= requested
&& channels.is_multiple_of(v)
&& operands.iter().all(|b| {
b.strides[..b.strides.len() - 1]
.iter()
.all(|&s| s.is_multiple_of(v))
})
})
.max()
.unwrap_or(1)
}
#[cfg(test)]
mod tests {
use super::*;
fn plan(width: usize, chans: usize) -> DepthwiseSpace {
DepthwiseSpace {
b: 2,
oh: 56,
ow: 56,
c: 512,
rh: 5,
rw: 5,
rows: 4,
cols: 4,
tile_c: 32 * width * chans,
width,
plane_size: 32,
}
}
#[test]
fn the_depthwise_routine_states_three_levels() {
assert_eq!(
plan(1, 1).partitioning().table(&LABELS).to_string(),
[
" b × oh × ow × c × rh × rw b × oh × ow × c × rh × rw",
"",
" ◦ · × · × · × · × · × · 1 × 1 × 4 × 1 × 5 × 5",
" ▪ 32 units interleaved · × · × · × 32 × · × · 1 × 1 × 4 × 32 × 5 × 5",
" ▤ 4 planes a cube · × 4 × · × · × · × · 1 × 4 × 4 × 32 × 5 × 5",
" ▣ 6272 cubes 2 × 14 × 14 × 16 × · × · 2 × 56 × 56 × 512 × 5 × 5",
"",
" └─ count ────────────────┘ └─ tile ──────────────────┘",
]
.join("\n")
);
}
#[test]
fn a_unit_holding_several_channel_lines_walks_them() {
assert_eq!(
plan(4, 2).partitioning().table(&LABELS).to_string(),
[
" b × oh × ow × c × rh × rw b × oh × ow × c × rh × rw",
"",
" ◦ · × · × · × · × · × · 1 × 1 × 4 × 4 × 5 × 5",
" ▪ 32 units interleaved · × · × · × 32 × · × · 1 × 1 × 4 × 128 × 5 × 5",
" ↻ 2 steps · × · × · × 2 × · × · 1 × 1 × 4 × 256 × 5 × 5",
" ▤ 4 planes a cube · × 4 × · × · × · × · 1 × 4 × 4 × 256 × 5 × 5",
" ▣ 784 cubes 2 × 14 × 14 × 2 × · × · 2 × 56 × 56 × 512 × 5 × 5",
"",
" └─ count ────────────────┘ └─ tile ──────────────────┘",
]
.join("\n")
);
}
}