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}