cubek-convolution 0.3.0

CubeK: Convolution Kernels
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
//! Depthwise convolution, as a direct client of the tile DSL.
//!
//! Every accelerated convolution routine in this crate refuses `groups != 1`, so a depthwise
//! layer autotunes against a single surviving candidate and never reaches an accelerated path at
//! all. This is the missing one, and it is deliberately *not* built on the blueprint/routine
//! machinery the rest of `kernels::forward` uses: a depthwise convolution is not a contraction
//! over channels, so a stage hierarchy sized for one has nothing to size here.
//!
//! What it is instead: a dense convolution contracts input channels into output channels, so the
//! channel pair appears in the two operands and not in the accumulator. A depthwise one has no
//! such pairing — each channel carries its own filter and reaches exactly one output channel — so
//! a single channel axis appears in *all three* operands. That makes it a batch axis, and the
//! contraction is over the window taps alone. Written that way it is a space and a projection,
//! and `Tile::mma` is the whole body.
//!
//! Everything about the layout follows from that one axis. Channels stay innermost (NHWC) and are
//! what the units are spent on, so consecutive units read consecutive channels of one pixel and
//! the read coalesces; and they are what every operand is *lined* along, so one instruction moves
//! a unit's whole cell. A depthwise pass has too little arithmetic per byte to be anything but
//! bandwidth-bound, and both of those are what let it reach the bandwidth.

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};

/// The register block the leaf runs under.
///
/// 64 scalars is the register budget, which at four channels to a line is the same sixteen cells
/// every tiling here blocks into. The edge split earns its second copy of the walk because a
/// window this wide leaves most instances clear of the padded border, and they should not pay a
/// guard for the few that straddle it. Unit fan-out does not: the lines run along the channel,
/// not along `K`.
const REGISTER_BLOCK: RegisterBlock = RegisterBlock::new(64).split_edge();

// Output positions, the channel axis every operand shares, and the window taps.
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);

/// What a partitioning over these axes prints as ([`Partitioning::labelled`]): an [`Axis`] is an
/// index, and only the kernel that assigned it knows what it stands for.
///
/// Only tests print one today. It widens when a caller does.
#[cfg(test)]
const LABELS: [(Axis, &str); 6] = [
    (B, "b"),
    (OH, "oh"),
    (OW, "ow"),
    (C, "c"),
    (RH, "rh"),
    (RW, "rw"),
];

/// The space one launch runs over, in the terms the kernel builds it from: the problem's
/// extents and the tiling. The kernel's comptime argument, so the space the launch sizes its
/// grid from is the space the kernel walks.
#[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 {
    /// The space's axes and their extents, every one static.
    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),
        ]
    }

    /// Three levels, outermost first. The first separates the output across the launch grid: an
    /// all-`sequential` level would put the whole convolution in one instance, which is a
    /// correct kernel and a useless one. The next two separate what one cube took across the
    /// cube's own threads: rows go to planes, channels to units. The taps stay whole
    /// throughout: they are the contraction, and every tap of one output position accumulates
    /// into the same register.
    /// The levels, stated **from the leaf up, in counts**: one unit's channel line by the
    /// cube's columns by one row; the plane's units across channels, in turns; the channel lines
    /// a unit holds past its first, where the tile is wider than the plane; the cube's rows
    /// across its planes; and a cube per box, the taps whole. The channel axis takes `X` so
    /// that the fastest-moving cube index is the one memory is contiguous along.
    ///
    /// Round-robin across the units, so a unit holding several channel lines takes every
    /// `plane_size`-th rather than a contiguous run: a contiguous run puts a stride between
    /// what neighbouring units read and breaks the coalescing the whole NHWC layout is for.
    /// Columns stay whole: they are the register block, not a split. The walk over a unit's
    /// further lines is stated only where there are any — the old units level was that walk
    /// as well, its length found by dividing the tile.
    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())
    }

    /// The space with the levels that cut it: what the leaf and the overhangs are read off.
    pub fn partitioning(&self) -> Partitioning {
        Partitioning::new(self.space(), self.levels())
    }

    /// The grid this launch runs on: channels on `X`, columns on `Y`, rows and batches on `Z`,
    /// a plane per row of the cube.
    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),
        )
    }
}

/// `out[b, oh, ow, c] = Σ_{rh, rw} w[rh, rw, c] · input[b, oh*sh + rh*dh - ph, ow*sw + rw*dw - pw,
/// c]`
///
/// The same leaf the dense convolution runs. `C` being one of the accumulator's own axes is what
/// makes it batched rather than contracted; the leaf reads that off the spaces.
///
/// The filter is the *lhs* and the map the rhs, which is not arbitrary: the leaf serves the rhs in
/// lines along the accumulator's innermost axis, and that axis is the channel. The filter follows
/// it there (one filter value per channel of the cell), which is what a batched contraction
/// needs and what `V > 1` is.
///
/// Three levels: this cube's box of the output with the taps whole, this plane's row of it, then
/// this unit's channel lines, whose column block the leaf walks with the whole tap window at
/// each cell.
#[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 {
            // The plane's units, under its walk over channel lines where a unit holds several.
            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));
            }
        }
    }
}

/// How one cube's share of the output is shaped, and how the cube's threads divide it.
///
/// The three numbers are three different jobs, which is why they are not one "tile size":
///
/// - `rows` is the plane count. One plane per output row, so it is what fills the cube.
/// - `cols` is the accumulator block one unit keeps in registers. Every column of it re-reads
///   the same filter and overlapping input, so it is what amortises both.
/// - `chans` is how many channel *lines* one unit owns. The lines are distributed to the units
///   interleaved, so whatever this is, consecutive units still read consecutive channels and the
///   read coalesces.
/// - `lines` is how many channels one of those lines covers — the width every operand is served
///   in. It is the one knob that trades the two things a depthwise pass is limited by against
///   each other, which is why it is stated and not derived: a wider line is fewer instructions
///   per channel, and also more registers per unit and a wider channel tile, so fewer units with
///   anything to do when the block is narrow. It is a ceiling, not a demand — the launch drops to
///   what the buffers can actually be served in.
///
/// The cube's channel tile is `plane_size · lines · chans`.
#[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 {
    /// Four planes of one row each, four output columns per unit, scalar channels.
    ///
    /// Small on both spatial axes on purpose: the window overlaps, so a cube's halo is what it
    /// re-reads, but a *wide* tile is also what pushes its far corner past the padded border and
    /// costs every instance in it the guarded walk. Four each is where those two meet on the
    /// shapes an encoder actually runs. The line width is the knob worth deciding per problem,
    /// which [`for_problem`](Self::for_problem) is.
    fn default() -> Self {
        Self {
            rows: 4,
            cols: 4,
            chans: 1,
            lines: 1,
        }
    }
}

impl DepthwiseTiling {
    /// A window with at least this many taps re-reads enough of its halo to run out of
    /// instructions before it runs out of bandwidth, which is the only regime where a wider line
    /// pays for the registers it costs. A 5x5 window is the first one an encoder runs that
    /// reaches it.
    const INSTRUCTION_BOUND_TAPS: usize = 25;

    /// ...and the channel axis has to stay wide enough to fill the grid once a wide line has
    /// divided its parallelism. Below this many cube-widths of channels, widening starves the
    /// grid instead of the bus.
    const WIDE_BLOCK_UNIT_MULTIPLE: usize = 8;

    /// The tiling to run a problem of this shape under.
    ///
    /// Only [`lines`](Self::lines) is decided here, and it is close to a single question: is this
    /// convolution short of instructions or short of bandwidth? A wide line is four times fewer
    /// instructions per channel, and also four times the registers per unit and a four-times
    /// wider channel tile. A deep window over a wide block is instruction-bound and takes the
    /// trade; everything else is already reading memory as fast as the device will read it, and
    /// pays the registers for nothing.
    ///
    /// Both thresholds are where that turnover was observed rather than where a model of the
    /// hardware puts it, so this is a derivation a device is allowed to disagree with. The
    /// `depthwise` benchmark catalogue is the instrument that settles it: running its `Fixed`
    /// entries against `Routine` is what says whether this rule still picks the right line.
    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()
        }
    }

    /// Reject a degenerate tiling before its dimensions reach division and space construction.
    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)
    }

    /// The channel edge one cube owns, checked because every factor is public configuration or
    /// runtime hardware data.
    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)
    }

    /// The space this tiling implies for a problem of these extents, in the form the kernel
    /// builds it from.
    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,
        }
    }
}

/// The three tensors this routine moves, named because they are all `TensorBinding` and a
/// positional triple lets two of them be swapped without a word from the compiler.
///
/// NHWC maps, and Burn's `[out_channels, kh, kw, in_channels / groups]` filter — whose trailing
/// axis is 1 for a depthwise convolution.
pub struct DepthwiseTensors {
    pub input: TensorBinding,
    pub weight: TensorBinding,
    pub out: TensorBinding,
}

/// Which tiling to run a problem under.
///
/// [`Routine`](Self::Routine) is what ships; [`Fixed`](Self::Fixed) is what the benchmark
/// catalogue sweeps, and what a test uses to reach a tiling the rule would not pick.
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum DepthwiseStrategy {
    /// Decided from the problem, by [`DepthwiseTiling::for_problem`].
    Routine,
    /// Stated by the caller.
    Fixed(DepthwiseTiling),
}

/// Launch a depthwise convolution under `strategy`.
///
/// # Errors
///
/// [`ConvSetupError::NotDepthwise`] when the convolution is not one filter per channel — the
/// tuner reads that as "this candidate does not apply here", which is exactly what it is — and
/// [`ConvSetupError::InvalidConfig`] when a fixed tiling is degenerate. Returns
/// [`ConvSetupError::Unknown`] when re-laying the filter cannot be launched.
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;

    // The filter, re-laid so the channel is its innermost dim like every other operand's. It has
    // to be: the leaf serves a cell in lines along the channel, and one filter value broadcast
    // over a line would give every channel of it the first channel's filter. This is the one
    // copy the routine makes, and it is the smallest tensor in the problem by three orders of
    // magnitude — a 1632-channel 5x5 filter is 163 KB against 60 MB of map.
    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,
            },
        )?
    };

    // A tile that does not divide its axis leaves the last cube short, and a short cube's
    // terminal tile is still the full comptime size — so the cells past the end are addressed and
    // have to be guarded. Per axis, because the guard is real work per access and the axes that
    // need one are rarely the same: a 48x48 map divides evenly by any tile here while a
    // 24-channel block never fills one plane's width of channels.
    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;

    // Two gathered physical axes, one per spatial pair, each carrying its padding as the
    // projection's constant term. Batch and channel ride identity. The channel comes last in the
    // logical order because that is the axis the operand lines along, and the innermost logical
    // axis is the one a line covers.
    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),
        ],
    ))
    // A padded border can be represented by either the beginning padding in the projection or
    // the output binding extending far enough for the final window to overhang the input. Guard
    // from the actual first and last accessed coordinates so end-only padding is covered too.
    .boundaries(&[
        None,
        guard(check_h || ragged_oh),
        guard(check_w || ragged_ow),
        guard(ragged_c),
    ]);

    // Read in place rather than staged. A shared-memory stage costs a cooperative fill and a
    // sync per cube, and a deep block — many channels over few output positions — has too few
    // output positions to amortise either; only the widest spatial shapes have enough.
    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(())
}

/// The problem, in the terms the space is built from. NHWC throughout.
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 {
    /// Read the problem off the bindings themselves, so the shapes the space is built from are
    /// the shapes the kernel will address rather than a second copy that can disagree with them.
    ///
    /// # Errors
    ///
    /// [`ConvSetupError::NotDepthwise`] when the convolution is not one filter per channel. The
    /// tuner reads that as "this candidate does not apply here", which is exactly what it is.
    fn new(
        tensors: &DepthwiseTensors,
        args: ConvolutionArgs<2>,
        groups: usize,
    ) -> Result<Self, ConvSetupError> {
        // NHWC throughout: [batch, h, w, channels].
        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,
            // Burn hands weights as [out_channels, kh, kw, in_channels / groups]; depthwise makes
            // that last axis 1, so the filter is [c, kh, kw] with the channel outermost.
            rh: tensors.weight.shape[1],
            rw: tensors.weight.shape[2],
            stride: args.stride,
            padding: args.padding,
            dilation: args.dilation,
        })
    }

    /// How many taps one filter has.
    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],
        )
    }

    /// The filter as `[kh, kw, c]`, contiguous.
    ///
    /// Burn stores it `[c, kh, kw]`, which is the one layout this kernel cannot read: the channel
    /// has to be the innermost dim for a line to cover a cell's worth of filter. Permuting the
    /// binding's existing strides re-presents that logical tensor without assuming anything about
    /// its storage; `into_contiguous` is what makes the new layout physical.
    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]);
        // `[C, kh, kw, 1] -> [kh, kw, C]`. The omitted axis is singleton, so it contributes no
        // offset; every surviving axis must retain its actual stride for sliced/strided bindings.
        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
}

/// The plane width the channel tile is sized to.
///
/// `plane_size_max` deliberately, and it is only safe because this kernel issues no plane
/// instruction: the leaf is [`Instruction::Registers`], the taps contract into a register rather
/// than across units, and `planes()`/[`Coverage::PlaneUnits`] here distribute work rather than
/// cooperate. So the width is a coalescing decision, and a device honouring a narrower one still
/// gets every unit of the tile from a real thread — `Space::cube_dim` sizes the launch from the
/// same number.
///
/// The moment a plane reduction appears in this kernel that stops being true: wgpu reports a
/// range on AMD RDNA (32/64) and Intel (8/32), and a reduction sized to the max would cover a
/// fraction of its row on a device honouring the min.
fn plane_units(client: &Client) -> usize {
    client.properties().hardware.plane_size_max as usize
}

/// The boundary an axis needs, or `None` when every read along it is in bounds by construction.
fn guard(ragged: bool) -> Option<Boundary> {
    ragged.then_some(Boundary::Zero)
}

/// The widest line the channel axis can be served in across all three operands, up to what the
/// tiling asked for.
///
/// Every gate below the request is a fact about the buffers rather than a preference: the channel
/// must be the contiguous dim, and the width must divide the channel count, since a partial line
/// has no cell to be.
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::*;

    /// A `5x5` pass over a `56x56` map of 512 channels, four planes of one row, four output
    /// columns a unit, on a 32-unit plane. The channel tile is what the caller varies: it is
    /// `plane_size * width * chans`, and it alone decides whether a unit holds one channel line
    /// or several.
    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,
        }
    }

    /// The three levels as a table, leaf up: each row's tile is the row below it times the count
    /// beside it, so a tiling that stops dividing an axis where it meant to keep going shows up
    /// as a row that no longer multiplies out. The taps never divide: they are the contraction,
    /// and every one accumulates into the same register.
    #[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")
        );
    }

    /// A channel tile wider than one pass of the plane's lines adds a fourth level, which is the
    /// walk the kernel's `lines_below_the_units` branch runs: the units sit under it, and a unit
    /// takes every 32nd line rather than a contiguous run.
    #[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")
        );
    }
}