Skip to main content

cubek_convolution/kernels/forward/
depthwise.rs

1//! Depthwise convolution, as a direct client of the tile DSL.
2//!
3//! Every accelerated convolution routine in this crate refuses `groups != 1`, so a depthwise
4//! layer autotunes against a single surviving candidate and never reaches an accelerated path at
5//! all. This is the missing one, and it is deliberately *not* built on the blueprint/routine
6//! machinery the rest of `kernels::forward` uses: a depthwise convolution is not a contraction
7//! over channels, so a stage hierarchy sized for one has nothing to size here.
8//!
9//! What it is instead: a dense convolution contracts input channels into output channels, so the
10//! channel pair appears in the two operands and not in the accumulator. A depthwise one has no
11//! such pairing — each channel carries its own filter and reaches exactly one output channel — so
12//! a single channel axis appears in *all three* operands. That makes it a batch axis, and the
13//! contraction is over the window taps alone. Written that way it is a space and a projection,
14//! and `Tile::mma` is the whole body.
15//!
16//! Everything about the layout follows from that one axis. Channels stay innermost (NHWC) and are
17//! what the lanes are spent on, so consecutive lanes read consecutive channels of one pixel and
18//! the read coalesces; and they are what every operand is *lined* along, so one instruction moves
19//! a lane's whole cell. A depthwise pass has too little arithmetic per byte to be anything but
20//! bandwidth-bound, and both of those are what let it reach the bandwidth.
21
22use cubecl::{
23    prelude::*,
24    server::LaunchError,
25    zspace::{Shape, Strides},
26};
27use cubek_std::InputBinding;
28use cubek_tile::*;
29
30use crate::{components::ConvSetupError, launch::ConvolutionArgs};
31
32/// The register block the leaf runs under.
33///
34/// 64 scalars is the register budget, which at four channels to a line is the same sixteen cells
35/// every tiling here blocks into. The edge split earns its second copy of the walk because a
36/// window this wide leaves most instances clear of the padded border, and they should not pay a
37/// guard for the few that straddle it. Lane fan-out does not: the lines run along the channel,
38/// not along `K`.
39const REGISTER_BLOCK: RegisterBlock = RegisterBlock::new(64).split_edge();
40
41// Output positions, the channel axis every operand shares, and the window taps.
42const B: Axis = Axis(0);
43const OH: Axis = Axis(1);
44const OW: Axis = Axis(2);
45const C: Axis = Axis(3);
46const RH: Axis = Axis(4);
47const RW: Axis = Axis(5);
48
49/// What a partitioning over these axes prints as ([`Partitioning::labelled`]): an [`Axis`] is an
50/// index, and only the kernel that assigned it knows what it stands for.
51///
52/// Only tests print one today. It widens when a caller does.
53#[cfg(test)]
54const LABELS: [(Axis, &str); 6] = [
55    (B, "b"),
56    (OH, "oh"),
57    (OW, "ow"),
58    (C, "c"),
59    (RH, "rh"),
60    (RW, "rw"),
61];
62
63/// The space one launch runs over, in the terms the kernel builds it from: the problem's
64/// extents and the tiling. The kernel's comptime argument, so the space the launch sizes its
65/// grid from is the space the kernel walks.
66#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
67pub struct DepthwiseSpace {
68    b: usize,
69    oh: usize,
70    ow: usize,
71    c: usize,
72    rh: usize,
73    rw: usize,
74    rows: usize,
75    cols: usize,
76    tile_c: usize,
77    width: usize,
78    plane_size: usize,
79}
80
81impl DepthwiseSpace {
82    /// The space's axes and their extents, every one static.
83    pub fn extents(&self) -> Vec<(Axis, usize)> {
84        vec![
85            (B, self.b),
86            (OH, self.oh),
87            (OW, self.ow),
88            (C, self.c),
89            (RH, self.rh),
90            (RW, self.rw),
91        ]
92    }
93
94    /// Three levels, outermost first. The first separates the output across the launch grid: an
95    /// all-`sequential` level would put the whole convolution in one instance, which is a
96    /// correct kernel and a useless one. The next two separate what one cube took across the
97    /// cube's own threads: rows go to planes, channels to lanes. The taps stay whole
98    /// throughout: they are the contraction, and every tap of one output position accumulates
99    /// into the same register.
100    /// The levels, stated **from the leaf up, in counts**: one lane's channel line by the
101    /// cube's columns by one row; the plane's lanes across channels, in turns; the channel lines
102    /// a lane holds past its first, where the tile is wider than the plane; the cube's rows
103    /// across its planes; and a cube per box, the taps whole. The channel axis takes `X` so
104    /// that the fastest-moving cube index is the one memory is contiguous along.
105    ///
106    /// Round-robin across the lanes, so a lane holding several channel lines takes every
107    /// `plane_size`-th rather than a contiguous run: a contiguous run puts a stride between
108    /// what neighbouring lanes read and breaks the coalescing the whole NHWC layout is for.
109    /// Columns stay whole: they are the register block, not a split. The walk over a lane's
110    /// further lines is stated only where there are any — the old lanes level was that walk
111    /// as well, its length found by dividing the tile.
112    pub fn levels(&self) -> Vec<Level> {
113        let Self {
114            rows,
115            cols,
116            tile_c,
117            width,
118            plane_size,
119            ..
120        } = *self;
121        let plane_c = width * plane_size;
122        assert!(
123            tile_c.is_multiple_of(plane_c),
124            "DepthwiseSpace: {plane_size} lanes of {width} channels do not divide a tile of {tile_c}"
125        );
126        let lanes = Tiling::leaf(&[(C, width), (OW, cols), (OH, 1)])
127            .lanes(&[(C, plane_size)])
128            .interleaved(C);
129        let lines = match tile_c / plane_c {
130            1 => lanes,
131            further => lanes.walk(&[(C, further)]),
132        };
133        lines
134            .planes(&[(OH, rows)])
135            .cubes(&[C, OW, OH])
136            .batches(&[B])
137            .levels()
138    }
139
140    pub fn space(&self) -> Space {
141        Space::new(&self.extents())
142    }
143
144    /// The space with the levels that cut it: what the leaf and the overhangs are read off.
145    pub fn partitioning(&self) -> Partitioning {
146        Partitioning::new(self.space(), self.levels())
147    }
148
149    /// The grid this launch runs on: channels on `X`, columns on `Y`, rows and batches on `Z`,
150    /// a plane per row of the cube.
151    pub fn grid(&self) -> (CubeCount, CubeDim) {
152        (
153            CubeCount::Static(
154                self.c.div_ceil(self.tile_c) as u32,
155                self.ow.div_ceil(self.cols) as u32,
156                (self.oh.div_ceil(self.rows) * self.b) as u32,
157            ),
158            CubeDim::new_2d(self.plane_size as u32, self.rows as u32),
159        )
160    }
161}
162
163/// `out[b, oh, ow, c] = Σ_{rh, rw} w[rh, rw, c] · input[b, oh*sh + rh*dh - ph, ow*sw + rw*dw - pw, c]`
164///
165/// The same leaf the dense convolution runs. `C` being one of the accumulator's own axes is what
166/// makes it batched rather than contracted; the leaf reads that off the spaces.
167///
168/// The filter is the *lhs* and the map the rhs, which is not arbitrary: the leaf serves the rhs in
169/// lines along the accumulator's innermost axis, and that axis is the channel. The filter follows
170/// it there (one filter value per channel of the cell), which is what a batched contraction
171/// needs and what `V > 1` is.
172///
173/// Three levels: this cube's box of the output with the taps whole, this plane's row of it, then
174/// this lane's channel lines, whose column block the leaf walks with the whole tap window at
175/// each cell.
176#[cube(launch)]
177fn depthwise_kernel<E: Numeric, V: Size>(
178    weight: &TileArg<'_, E, V>,
179    input: &TileArg<'_, E, V>,
180    out: &TileArg<'_, E, V>,
181    space: Partitioning,
182    #[define(E)] _dtype: ElemType,
183) {
184    let weight = weight.tile(comptime!(space.clone()));
185    let input = input.tile(comptime!(space.clone()));
186    let out = out.tile(comptime!(space.clone()));
187
188    // Three levels, or four where a lane holds channel lines past its first: then the plane's
189    // walk is over those lines and the lanes sit under it.
190    let lines_below_the_lanes = comptime!(space.levels().len() > 3);
191    for cube in space {
192        let out = out.at(&cube);
193        let weight = weight.at(&cube);
194        let input = input.at(&cube);
195        for plane in cube {
196            for step in plane {
197                if lines_below_the_lanes {
198                    for lane in step {
199                        let mut out = out.at(&lane);
200                        out.mm_with(
201                            &weight.at(&lane),
202                            &input.at(&lane),
203                            REGISTER_BLOCK,
204                            Semiring::SUM_PROD,
205                        );
206                    }
207                } else {
208                    let lane = step;
209                    let mut out = out.at(&lane);
210                    out.mm_with(
211                        &weight.at(&lane),
212                        &input.at(&lane),
213                        REGISTER_BLOCK,
214                        Semiring::SUM_PROD,
215                    );
216                }
217            }
218        }
219    }
220}
221
222/// How one cube's share of the output is shaped, and how the cube's threads divide it.
223///
224/// The three numbers are three different jobs, which is why they are not one "tile size":
225///
226/// - `rows` is the plane count. One plane per output row, so it is what fills the cube.
227/// - `cols` is the accumulator block one lane keeps in registers. Every column of it re-reads
228///   the same filter and overlapping input, so it is what amortises both.
229/// - `chans` is how many channel *lines* one lane owns. Lanes are dealt lines interleaved, so
230///   whatever this is, consecutive lanes still read consecutive channels and the read coalesces.
231/// - `lines` is how many channels one of those lines covers — the width every operand is served
232///   in. It is the one knob that trades the two things a depthwise pass is limited by against
233///   each other, which is why it is stated and not derived: a wider line is fewer instructions
234///   per channel, and also more registers per lane and a wider channel tile, so fewer lanes with
235///   anything to do when the block is narrow. It is a ceiling, not a demand — the launch drops to
236///   what the buffers can actually be served in.
237///
238/// The cube's channel tile is `plane_size · lines · chans`.
239#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
240pub struct DepthwiseTiling {
241    pub rows: usize,
242    pub cols: usize,
243    pub chans: usize,
244    pub lines: usize,
245}
246
247impl Default for DepthwiseTiling {
248    /// Four planes of one row each, four output columns per lane, scalar channels.
249    ///
250    /// Small on both spatial axes on purpose: the window overlaps, so a cube's halo is what it
251    /// re-reads, but a *wide* tile is also what pushes its far corner past the padded border and
252    /// costs every instance in it the guarded walk. Four each is where those two meet on the
253    /// shapes an encoder actually runs. The line width is the knob worth deciding per problem,
254    /// which [`for_problem`](Self::for_problem) is.
255    fn default() -> Self {
256        Self {
257            rows: 4,
258            cols: 4,
259            chans: 1,
260            lines: 1,
261        }
262    }
263}
264
265impl DepthwiseTiling {
266    /// A window with at least this many taps re-reads enough of its halo to run out of
267    /// instructions before it runs out of bandwidth, which is the only regime where a wider line
268    /// pays for the registers it costs. A 5x5 window is the first one an encoder runs that
269    /// reaches it.
270    const INSTRUCTION_BOUND_TAPS: usize = 25;
271
272    /// ...and the channel axis has to stay wide enough to fill the grid once a wide line has
273    /// divided its parallelism. Below this many cube-widths of channels, widening starves the
274    /// grid instead of the bus.
275    const WIDE_BLOCK_LANE_MULTIPLE: usize = 8;
276
277    /// The tiling to run a problem of this shape under.
278    ///
279    /// Only [`lines`](Self::lines) is decided here, and it is close to a single question: is this
280    /// convolution short of instructions or short of bandwidth? A wide line is four times fewer
281    /// instructions per channel, and also four times the registers per lane and a four-times
282    /// wider channel tile. A deep window over a wide block is instruction-bound and takes the
283    /// trade; everything else is already reading memory as fast as the device will read it, and
284    /// pays the registers for nothing.
285    ///
286    /// Both thresholds are where that turnover was observed rather than where a model of the
287    /// hardware puts it, so this is a derivation a device is allowed to disagree with. The
288    /// `depthwise` benchmark catalogue is the instrument that settles it: running its `Fixed`
289    /// entries against `Routine` is what says whether this rule still picks the right line.
290    pub fn for_problem(channels: usize, taps: usize, lanes: usize) -> Self {
291        let deep_window = taps >= Self::INSTRUCTION_BOUND_TAPS;
292        let wide_block = channels >= Self::WIDE_BLOCK_LANE_MULTIPLE * lanes;
293
294        Self {
295            lines: match deep_window && wide_block {
296                true => 4,
297                false => 1,
298            },
299            ..Default::default()
300        }
301    }
302
303    /// Reject a degenerate tiling before its dimensions reach division and space construction.
304    fn validate(self) -> Result<Self, ConvSetupError> {
305        if self.rows == 0 || self.cols == 0 || self.chans == 0 || self.lines == 0 {
306            return Err(ConvSetupError::InvalidConfig(Box::new(format!(
307                "depthwise tiling dimensions must be non-zero, got rows {}, cols {}, chans {}, \
308                 lines {}",
309                self.rows, self.cols, self.chans, self.lines
310            ))));
311        }
312        Ok(self)
313    }
314
315    /// The channel edge one cube owns, checked because every factor is public configuration or
316    /// runtime hardware data.
317    fn channel_tile(self, lanes: usize, width: usize) -> Result<usize, ConvSetupError> {
318        let tile = lanes
319            .checked_mul(width)
320            .and_then(|tile| tile.checked_mul(self.chans))
321            .ok_or_else(|| {
322                ConvSetupError::InvalidConfig(Box::new(format!(
323                    "depthwise channel tile overflows: {lanes} lanes * {width} channels/line * {} \
324                     lines/lane",
325                    self.chans
326                )))
327            })?;
328        if tile == 0 {
329            return Err(ConvSetupError::InvalidConfig(Box::new(format!(
330                "depthwise channel tile must be non-zero, got {lanes} lanes * {width} \
331                 channels/line * {} lines/lane",
332                self.chans
333            ))));
334        }
335        Ok(tile)
336    }
337
338    /// The space this tiling implies for a problem of these extents, in the form the kernel
339    /// builds it from.
340    fn plan(
341        &self,
342        geometry: &Geometry,
343        lanes: usize,
344        tile_c: usize,
345        width: usize,
346    ) -> DepthwiseSpace {
347        DepthwiseSpace {
348            b: geometry.b,
349            oh: geometry.oh,
350            ow: geometry.ow,
351            c: geometry.c,
352            rh: geometry.rh,
353            rw: geometry.rw,
354            rows: self.rows,
355            cols: self.cols,
356            tile_c,
357            width,
358            plane_size: lanes,
359        }
360    }
361}
362
363/// The three tensors this routine moves, named because they are all `TensorBinding` and a
364/// positional triple lets two of them be swapped without a word from the compiler.
365///
366/// NHWC maps, and Burn's `[out_channels, kh, kw, in_channels / groups]` filter — whose trailing
367/// axis is 1 for a depthwise convolution.
368pub struct DepthwiseTensors {
369    pub input: TensorBinding,
370    pub weight: TensorBinding,
371    pub out: TensorBinding,
372}
373
374/// Which tiling to run a problem under.
375///
376/// [`Routine`](Self::Routine) is what ships; [`Fixed`](Self::Fixed) is what the benchmark
377/// catalogue sweeps, and what a test uses to reach a tiling the rule would not pick.
378#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
379pub enum DepthwiseStrategy {
380    /// Decided from the problem, by [`DepthwiseTiling::for_problem`].
381    Routine,
382    /// Stated by the caller.
383    Fixed(DepthwiseTiling),
384}
385
386/// Launch a depthwise convolution under `strategy`.
387///
388/// # Errors
389///
390/// [`ConvSetupError::NotDepthwise`] when the convolution is not one filter per channel — the
391/// tuner reads that as "this candidate does not apply here", which is exactly what it is — and
392/// [`ConvSetupError::InvalidConfig`] when a fixed tiling is degenerate. Returns
393/// [`ConvSetupError::Unknown`] when re-laying the filter cannot be launched.
394pub fn launch_depthwise(
395    client: &Client,
396    tensors: DepthwiseTensors,
397    args: ConvolutionArgs<2>,
398    groups: usize,
399    dtype: ElemType,
400    strategy: DepthwiseStrategy,
401) -> Result<(), ConvSetupError> {
402    let geometry = Geometry::new(&tensors, args, groups)?;
403    let lanes = plane_lanes(client);
404    let tiling = match strategy {
405        DepthwiseStrategy::Routine => {
406            DepthwiseTiling::for_problem(geometry.c, geometry.taps(), lanes)
407        }
408        DepthwiseStrategy::Fixed(tiling) => tiling,
409    }
410    .validate()?;
411    let DepthwiseTensors { input, weight, out } = tensors;
412
413    // The filter, re-laid so the channel is its innermost dim like every other operand's. It has
414    // to be: the leaf serves a cell in lines along the channel, and one filter value broadcast
415    // over a line would give every channel of it the first channel's filter. This is the one
416    // copy the routine makes, and it is the smallest tensor in the problem by three orders of
417    // magnitude — a 1632-channel 5x5 filter is 163 KB against 60 MB of map.
418    let weight = geometry
419        .channels_innermost(client, weight, dtype)
420        .map_err(|_| ConvSetupError::Unknown)?;
421
422    let width = line_width(
423        client,
424        geometry.c,
425        dtype,
426        tiling.lines,
427        &[&input, &weight, &out],
428    );
429    let tile_c = tiling.channel_tile(lanes, width)?;
430    let plan = tiling.plan(&geometry, lanes, tile_c, width);
431    let launch =
432        Launcher::partitioned(client, plan.partitioning(), plan.grid(), KernelForm::Static);
433
434    // A tile that does not divide its axis leaves the last cube short, and a short cube's
435    // terminal tile is still the full comptime size — so the cells past the end are addressed and
436    // have to be guarded. Per axis, because the guard is real work per access and the axes that
437    // need one are rarely the same: a 48x48 map divides evenly by any tile here while a
438    // 24-channel block never fills one lane-width.
439    let ragged_c = !geometry.c.is_multiple_of(tile_c);
440    let ragged_oh = !geometry.oh.is_multiple_of(tiling.rows);
441    let ragged_ow = !geometry.ow.is_multiple_of(tiling.cols);
442    let check_h = geometry.should_check_height_bounds();
443    let check_w = geometry.should_check_width_bounds();
444    let [ph, pw] = geometry.padding;
445    let [sh, sw] = geometry.stride;
446    let [dh, dw] = geometry.dilation;
447
448    // Two gathered physical axes, one per spatial pair, each carrying its padding as the
449    // projection's constant term. Batch and channel ride identity. The channel comes last in the
450    // logical order because that is the axis the operand lines along, and the innermost logical
451    // axis is the one a line covers.
452    let in_spec = TileSpec::new(Projection::new(
453        &[B, OH, OW, RH, RW, C],
454        &[
455            PhysicalAxisMap::of(B),
456            PhysicalAxisMap::affine_with_offset(&[(OH, sh), (RH, dh)], -(ph as isize)),
457            PhysicalAxisMap::affine_with_offset(&[(OW, sw), (RW, dw)], -(pw as isize)),
458            PhysicalAxisMap::of(C),
459        ],
460    ))
461    // A padded border can be represented by either the beginning padding in the projection or
462    // the output binding extending far enough for the final window to overhang the input. Guard
463    // from the actual first and last accessed coordinates so end-only padding is covered too.
464    .boundaries(&[
465        None,
466        guard(check_h || ragged_oh),
467        guard(check_w || ragged_ow),
468        guard(ragged_c),
469    ]);
470
471    // Read in place rather than staged. A shared-memory stage costs a cooperative fill and a
472    // sync per cube, and a deep block — many channels over few output positions — has too few
473    // output positions to amortise either; only the widest spatial shapes have enough.
474    let w_spec = TileSpec::direct(&[RH, RW, C]).boundaries(&[None, None, guard(ragged_c)]);
475    let out_spec = TileSpec::direct(&[B, OH, OW, C]).boundaries(&[
476        None,
477        guard(ragged_oh),
478        guard(ragged_ow),
479        guard(ragged_c),
480    ]);
481
482    depthwise_kernel::launch(
483        client,
484        launch.cube_count(),
485        launch.cube_dim(),
486        width,
487        TileArgLaunch::new(weight.into_tensor_arg(), w_spec),
488        TileArgLaunch::new(input.into_tensor_arg(), in_spec),
489        TileArgLaunch::new(out.into_tensor_arg(), out_spec),
490        launch.partitioning_arg(),
491        dtype,
492    );
493
494    Ok(())
495}
496
497/// The problem, in the terms the space is built from. NHWC throughout.
498struct Geometry {
499    b: usize,
500    ih: usize,
501    iw: usize,
502    oh: usize,
503    ow: usize,
504    c: usize,
505    rh: usize,
506    rw: usize,
507    stride: [usize; 2],
508    padding: [usize; 2],
509    dilation: [usize; 2],
510}
511
512impl Geometry {
513    /// Read the problem off the bindings themselves, so the shapes the space is built from are
514    /// the shapes the kernel will address rather than a second copy that can disagree with them.
515    ///
516    /// # Errors
517    ///
518    /// [`ConvSetupError::NotDepthwise`] when the convolution is not one filter per channel. The
519    /// tuner reads that as "this candidate does not apply here", which is exactly what it is.
520    fn new(
521        tensors: &DepthwiseTensors,
522        args: ConvolutionArgs<2>,
523        groups: usize,
524    ) -> Result<Self, ConvSetupError> {
525        // NHWC throughout: [batch, h, w, channels].
526        let input_channels = tensors.input.shape[3];
527        let output_channels = tensors.out.shape[3];
528        let weight_channels = tensors.weight.shape[0];
529        let weight_group_channels = tensors.weight.shape[3];
530        if groups != input_channels
531            || output_channels != input_channels
532            || weight_channels != input_channels
533            || weight_group_channels != 1
534        {
535            return Err(ConvSetupError::NotDepthwise {
536                groups,
537                input_channels,
538                output_channels,
539                weight_channels,
540                weight_group_channels,
541            });
542        }
543
544        Ok(Self {
545            b: tensors.out.shape[0],
546            ih: tensors.input.shape[1],
547            iw: tensors.input.shape[2],
548            oh: tensors.out.shape[1],
549            ow: tensors.out.shape[2],
550            c: input_channels,
551            // Burn hands weights as [out_channels, kh, kw, in_channels / groups]; depthwise makes
552            // that last axis 1, so the filter is [c, kh, kw] with the channel outermost.
553            rh: tensors.weight.shape[1],
554            rw: tensors.weight.shape[2],
555            stride: args.stride,
556            padding: args.padding,
557            dilation: args.dilation,
558        })
559    }
560
561    /// How many taps one filter has.
562    fn taps(&self) -> usize {
563        self.rh * self.rw
564    }
565
566    fn should_check_height_bounds(&self) -> bool {
567        spatial_bounds_required(
568            self.ih,
569            self.oh,
570            self.rh,
571            self.stride[0],
572            self.padding[0],
573            self.dilation[0],
574        )
575    }
576
577    fn should_check_width_bounds(&self) -> bool {
578        spatial_bounds_required(
579            self.iw,
580            self.ow,
581            self.rw,
582            self.stride[1],
583            self.padding[1],
584            self.dilation[1],
585        )
586    }
587
588    /// The filter as `[kh, kw, c]`, contiguous.
589    ///
590    /// Burn stores it `[c, kh, kw]`, which is the one layout this kernel cannot read: the channel
591    /// has to be the innermost dim for a line to cover a cell's worth of filter. Permuting the
592    /// binding's existing strides re-presents that logical tensor without assuming anything about
593    /// its storage; `into_contiguous` is what makes the new layout physical.
594    fn channels_innermost(
595        &self,
596        client: &Client,
597        weight: TensorBinding,
598        dtype: ElemType,
599    ) -> Result<TensorBinding, LaunchError> {
600        let mut permuted = weight;
601        let channel_stride = permuted.strides[0];
602        let row_stride = permuted.strides[1];
603        let col_stride = permuted.strides[2];
604        permuted.shape = Shape::from(vec![self.rh, self.rw, self.c]);
605        // `[C, kh, kw, 1] -> [kh, kw, C]`. The omitted axis is singleton, so it contributes no
606        // offset; every surviving axis must retain its actual stride for sliced/strided bindings.
607        permuted.strides = Strides::new(&[row_stride, col_stride, channel_stride]);
608
609        Ok(InputBinding::new(permuted, dtype)
610            .into_contiguous(client)?
611            .into_data())
612    }
613}
614
615fn spatial_bounds_required(
616    input_size: usize,
617    output_size: usize,
618    kernel_size: usize,
619    stride: usize,
620    padding: usize,
621    dilation: usize,
622) -> bool {
623    let first = -(padding as i64);
624    let last = (output_size as i64 - 1) * stride as i64
625        + (kernel_size as i64 - 1) * dilation as i64
626        - padding as i64;
627
628    first < 0 || last >= input_size as i64
629}
630
631/// The plane width the channel tile is sized to.
632///
633/// `plane_size_max` deliberately, and it is only safe because this kernel issues no plane
634/// instruction: the leaf is [`Instruction::Registers`], the taps contract into a register rather
635/// than across lanes, and `planes()`/[`Coverage::PlaneLanes`] here distribute work rather than
636/// cooperate. So the width is a coalescing decision, and a device honouring a narrower one still
637/// gets every lane of the tile from a real thread — `Space::cube_dim` sizes the launch from the
638/// same number.
639///
640/// The moment a plane reduction appears in this kernel that stops being true: wgpu reports a
641/// range on AMD RDNA (32/64) and Intel (8/32), and a reduction sized to the max would cover a
642/// fraction of its row on a device honouring the min.
643fn plane_lanes(client: &Client) -> usize {
644    client.properties().hardware.plane_size_max as usize
645}
646
647/// The boundary an axis needs, or `None` when every read along it is in bounds by construction.
648fn guard(ragged: bool) -> Option<Boundary> {
649    ragged.then_some(Boundary::Zero)
650}
651
652/// The widest line the channel axis can be served in across all three operands, up to what the
653/// tiling asked for.
654///
655/// Every gate below the request is a fact about the buffers rather than a preference: the channel
656/// must be the contiguous dim, and the width must divide the channel count, since a partial line
657/// has no cell to be.
658fn line_width(
659    client: &Client,
660    channels: usize,
661    dtype: ElemType,
662    requested: usize,
663    operands: &[&TensorBinding],
664) -> usize {
665    if !operands.iter().all(|b| b.strides.last() == Some(&1)) {
666        return 1;
667    }
668
669    client
670        .io_optimized_vector_sizes(dtype.size())
671        .filter(|&v| {
672            v <= requested
673                && channels.is_multiple_of(v)
674                && operands.iter().all(|b| {
675                    b.strides[..b.strides.len() - 1]
676                        .iter()
677                        .all(|&s| s.is_multiple_of(v))
678                })
679        })
680        .max()
681        .unwrap_or(1)
682}
683
684#[cfg(test)]
685mod tests {
686    use super::*;
687
688    /// A `5x5` pass over a `56x56` map of 512 channels, four planes of one row, four output
689    /// columns a lane, on a 32-lane plane. The channel tile is what the caller varies: it is
690    /// `plane_size * width * chans`, and it alone decides whether a lane holds one channel line
691    /// or several.
692    fn plan(width: usize, chans: usize) -> DepthwiseSpace {
693        DepthwiseSpace {
694            b: 2,
695            oh: 56,
696            ow: 56,
697            c: 512,
698            rh: 5,
699            rw: 5,
700            rows: 4,
701            cols: 4,
702            tile_c: 32 * width * chans,
703            width,
704            plane_size: 32,
705        }
706    }
707
708    /// The three levels as a table, leaf up: each row's tile is the row below it times the count
709    /// beside it, so a tiling that stops dividing an axis where it meant to keep going shows up
710    /// as a row that no longer multiplies out. The taps never divide: they are the contraction,
711    /// and every one accumulates into the same register.
712    #[test]
713    fn the_depthwise_routine_states_three_levels() {
714        assert_eq!(
715            plan(1, 1).partitioning().labelled(&LABELS).to_string(),
716            [
717                "        b × oh × ow ×  c × rh × rw    b × oh × ow ×   c × rh × rw",
718                "",
719                "  ◦     · ×  · ×  · ×  · ×  · ×  ·    1 ×  1 ×  4 ×   1 ×  5 ×  5",
720                "  ▪     · ×  · ×  · × 32 ×  · ×  ·    1 ×  1 ×  4 ×  32 ×  5 ×  5",
721                "  ▤     · ×  4 ×  · ×  · ×  · ×  ·    1 ×  4 ×  4 ×  32 ×  5 ×  5",
722                "  ▣     2 × 14 × 14 × 16 ×  · ×  ·    2 × 56 × 56 × 512 ×  5 ×  5",
723                "",
724                "        └─ count ────────────────┘    └─ tile ──────────────────┘",
725            ]
726            .join("\n")
727        );
728    }
729
730    /// A channel tile wider than one pass of the plane's lines adds a fourth level, which is the
731    /// walk the kernel's `lines_below_the_lanes` branch runs: the lanes sit under it, and a lane
732    /// takes every 32nd line rather than a contiguous run.
733    #[test]
734    fn a_lane_holding_several_channel_lines_walks_them() {
735        assert_eq!(
736            plan(4, 2).partitioning().labelled(&LABELS).to_string(),
737            [
738                "        b × oh × ow ×  c × rh × rw    b × oh × ow ×   c × rh × rw",
739                "",
740                "  ◦     · ×  · ×  · ×  · ×  · ×  ·    1 ×  1 ×  4 ×   4 ×  5 ×  5",
741                "  ▪     · ×  · ×  · × 32 ×  · ×  ·    1 ×  1 ×  4 × 128 ×  5 ×  5",
742                "  ↻     · ×  · ×  · ×  2 ×  · ×  ·    1 ×  1 ×  4 × 256 ×  5 ×  5",
743                "  ▤     · ×  4 ×  · ×  · ×  · ×  ·    1 ×  4 ×  4 × 256 ×  5 ×  5",
744                "  ▣     2 × 14 × 14 ×  2 ×  · ×  ·    2 × 56 × 56 × 512 ×  5 ×  5",
745                "",
746                "        └─ count ────────────────┘    └─ tile ──────────────────┘",
747            ]
748            .join("\n")
749        );
750    }
751}