burn-cubecl 0.22.0-pre.4

Generic backend that can be compiled just-in-time to any shader language target
Documentation
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
use burn_backend::cubecl::dtype_to_storage_type;
use burn_backend::{
    DType,
    ops::{ConvOptions, conv::calculate_conv_output_sizes},
};
use burn_std::{Metadata, Shape, Slice};
use core::iter;
use cubecl::{
    prelude::*,
    std::tensor::{TensorHandle, into_contiguous_pitched},
};
use cubek::convolution::components::ConvSetupError;

use crate::{
    CubeDevice,
    kernel::{
        AddOp, into_contiguous_aligned, launch_binop,
        matmul::{MatmulStrategy, matmul},
        reduce::{KernelReduceStrategy, reduce_dim},
        slice_assign, slice_with_steps,
        utils::split_dim,
    },
    ops::{
        numeric::{empty_device_dtype, zeros_client},
        reshape, swap_dims,
    },
    tensor::CubeTensor,
};
use cubek::reduce::components::instructions::ReduceOperationConfig;

#[cfg(not(test))]
pub(crate) fn batches_per_run(
    batch_size: usize,
    out_shape: usize,
    plane_size: usize,
) -> Result<usize, ConvSetupError> {
    use cubek::matmul::definition::MatmulAvailabilityError;

    let cube_count_per_batch = out_shape.div_ceil(plane_size);
    let max_cube_count = u16::MAX as usize;
    let max_simultaneous = Ord::min(max_cube_count / cube_count_per_batch, batch_size);
    if max_simultaneous == 0 {
        return Err(MatmulAvailabilityError::CubeCountTooBig(CubeCount::Static(
            cube_count_per_batch as u32,
            1,
            1,
        ))
        .into());
    }
    Ok((0..=max_simultaneous)
        .rev()
        .find(|per_run| batch_size.is_multiple_of(*per_run))
        .expect("Logically not possible"))
}

#[cfg(test)]
#[allow(unused)]
pub(crate) fn batches_per_run(
    batch_size: usize,
    out_shape: usize,
    plane_size: usize,
) -> Result<usize, ConvSetupError> {
    Ok(1)
}

pub fn conv_im2col_1x1<const N: usize>(
    input: CubeTensor,
    weight: CubeTensor,
    bias: Option<CubeTensor>,
    options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
    let rank = input.meta.num_dims();
    let dim_c = rank - 1;

    let out_channels = weight.meta.shape()[0];

    check_pointwise_strided(&weight.meta.shape()[1..dim_c], &options)?;

    let out_shape = calculate_conv_output_sizes(
        &weight.meta.shape()[1..dim_c],
        &options.stride,
        &options.padding,
        &options.dilation,
        &input.meta.shape()[1..dim_c],
    );

    let mut split_m = vec![input.meta.shape()[0]];
    split_m.extend(out_shape.iter().copied());

    let input = match options.stride.iter().all(|stride| *stride == 1) {
        true => input,
        false => strided_spatial_view(input, &out_shape, &options.stride),
    };

    let input = reshape_input(input); // [(NHW), C] : [M, K]
    let dtype = input.dtype;

    // Permute to N-major, while keeping memory layout K-major. K-major for both sides is the most
    // efficient for matmul, and allows skipping a contiguous kernel
    let weight = swap_dims(reshape_weight(weight), 0, 1); // [K, N]

    let out = matmul(input, weight, None, MatmulStrategy::default(), dtype)?; // [M, N]

    // Skip reshape to avoid potential `into_contiguous`. We're only splitting dims so it's safe.
    let mut out = split_dim(out, 0, &split_m); // [N, H, W, C]

    if let Some(bias) = bias {
        let mut bias_shape = iter::repeat_n(1, rank - 1).collect::<Vec<_>>();
        bias_shape.push(out_channels);
        let bias = reshape(bias, bias_shape.into());
        out = launch_binop::<AddOp>(out, bias);
    }

    Ok(out)
}

/// Reshapes NHWC input to [(N, H, W), C]
fn reshape_input(input: CubeTensor) -> CubeTensor {
    let input = crate::kernel::untile(input);
    let rank = input.meta.num_dims();
    let dim_c = rank - 1;
    let dtype = input.dtype;

    let batch_size = input.meta.shape()[0];
    let in_c: usize = input.meta.shape()[dim_c];
    let in_shape: Shape = input.meta.shape()[1..dim_c].into();

    let mut input = if !is_spatial_contiguous(input.meta.shape(), input.meta.strides()) {
        let (client, device) = (input.client.clone(), input.device.clone());
        let contiguous =
            into_contiguous_pitched(&client, input.binding(), dtype_to_storage_type(dtype));
        from_handle(client, device, contiguous, dtype)
    } else {
        input
    };

    *input.meta = Metadata::new(
        [batch_size * in_shape.num_elements(), in_c], // [M, K]
        [input.meta.strides()[dim_c - 1], input.meta.strides()[dim_c]],
    );
    input
}

fn is_spatial_contiguous(shape: &[usize], strides: &[usize]) -> bool {
    let rank = shape.len();
    let dim_c = rank - 1;

    // Channel must be contiguous for the [(N, H, W), C] reshape to be valid
    if strides[dim_c] != 1 {
        return false;
    }

    for i in (1..dim_c).rev() {
        if strides[i + 1] * shape[i + 1] != strides[i] {
            return false;
        }
    }
    true
}

fn from_handle(
    client: Client,
    device: CubeDevice,
    handle: TensorHandle,
    dtype: DType,
) -> CubeTensor {
    CubeTensor::new(
        client.clone(),
        handle.handle,
        *handle.metadata,
        device.clone(),
        dtype,
    )
}

/// Errors unless the convolution is pointwise, the case this module reduces to
/// a single matmul.
///
/// A 1x1 convolution with unit stride, no padding and no dilation maps every
/// output pixel to the input pixel under it, so `im2col` is the identity and
/// the convolution is a per-pixel `[C_in, C_out]` matmul. A 1x1 that strides or
/// pads reads outside its own pixel and is declined, even where its output
/// happens to come back the size of its input — `in = 2 * padding + 1` under a
/// stride of 2 is such a shape.
///
/// The shapes are NHWC, as everything below `conv/base.rs` is.
fn check_pointwise<const N: usize>(
    kernel_shape: &[usize],
    options: &ConvOptions<N>,
) -> Result<(), ConvSetupError> {
    check_pointwise_strided(kernel_shape, options)?;

    match options.stride.iter().all(|stride| *stride == 1) {
        true => Ok(()),
        false => Err(ConvSetupError::Unknown),
    }
}

/// The same, minus the stride: a strided pointwise convolution reads a regular
/// subset of its input, which a view can hold exactly, so the forward pass
/// still lowers to one matmul. The gradient paths have no such view and keep
/// [`check_pointwise`].
fn check_pointwise_strided<const N: usize>(
    kernel_shape: &[usize],
    options: &ConvOptions<N>,
) -> Result<(), ConvSetupError> {
    if options.groups != 1 {
        return Err(ConvSetupError::Groups(options.groups));
    }

    let pointwise = kernel_shape.iter().all(|size| *size == 1)
        && options
            .padding
            .iter()
            .all(|&(begin, end)| begin == 0 && end == 0)
        && options.dilation.iter().all(|dilation| *dilation == 1);

    match pointwise {
        true => Ok(()),
        false => Err(ConvSetupError::Unknown),
    }
}

/// The view a strided pointwise convolution actually reads: every `stride`-th
/// position along each spatial dim. Multiplying the spatial strides leaves the
/// contiguous copy in [`reshape_input`] to gather exactly those rows.
fn strided_spatial_view(input: CubeTensor, out_shape: &[usize], stride: &[usize]) -> CubeTensor {
    let mut input = crate::kernel::untile(input);
    let mut shape = input.meta.shape().to_vec();
    let mut strides = input.meta.strides().to_vec();

    for (dim, (out, step)) in out_shape.iter().zip(stride).enumerate() {
        shape[dim + 1] = *out;
        strides[dim + 1] *= *step;
    }

    *input.meta = Metadata::new(shape, strides);
    input
}

/// Drops a pointwise weight's unit kernel dimensions, giving `[C_out, C_in]`.
///
/// Rewriting the metadata rather than reshaping keeps a padded channel stride,
/// so a weight the pitched allocator already aligned for TMA is not copied to
/// say so. One that is not gets a pitched copy here rather than a second kernel
/// inside the matmul.
fn reshape_weight(weight: CubeTensor) -> CubeTensor {
    let mut weight = crate::kernel::untile(weight);
    let dim_c = weight.meta.num_dims() - 1;
    let strides = [weight.meta.strides()[0], weight.meta.strides()[dim_c]];
    let shape = [weight.meta.shape()[0], weight.meta.shape()[dim_c]];

    *weight.meta = Metadata::new(shape, strides);

    match strides[1] {
        1 => weight,
        _ => into_contiguous_aligned(weight),
    }
}

/// The gradient of a pointwise convolution with respect to its input, as one
/// matmul.
///
/// `grad_in[(n, h, w), c_in] = sum over c_out of grad_out[(n, h, w), c_out] *
/// weight[c_out, c_in]`. The fallback computes the same thing as a transposed
/// convolution, which has no NHWC path and falls back to a naive kernel on a
/// device with no accelerated matmul for the dtype.
pub fn dgrad_im2col_1x1<const N: usize>(
    out_grad: CubeTensor,
    weight: CubeTensor,
    input_shape: Shape,
    options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
    let dim_c = out_grad.meta.num_dims() - 1;

    check_pointwise(&weight.meta.shape()[1..dim_c], &options)?;

    let split_m = input_shape[..dim_c].to_vec();

    let out_grad = reshape_input(out_grad); // [(NHW), C_out] : [M, K]
    let dtype = out_grad.dtype;

    // No transpose here, unlike the forward: it wants `[K, N]` and the weight's
    // own order *is* `[C_out, C_in]`, because this reduces over `C_out` where
    // the forward reduces over `C_in`.
    let weight = reshape_weight(weight); // [K, N]

    let out = matmul(out_grad, weight, None, MatmulStrategy::default(), dtype)?; // [M, N]

    // Skip reshape to avoid potential `into_contiguous`. We're only splitting dims so it's safe.
    Ok(split_dim(out, 0, &split_m)) // [N, H, W, C_in]
}

/// The gradient of a pointwise convolution with respect to its weight, as one
/// matmul.
///
/// `grad_w[c_out, c_in] = sum over (n, h, w) of grad_out[(n, h, w), c_out] *
/// input[(n, h, w), c_in]` — a tall reduction into a small output, which the
/// fallback's "convolve the input by the gradient" framing hides from the
/// matmul tuner.
pub fn wgrad_im2col_1x1<const N: usize>(
    input: CubeTensor,
    out_grad: CubeTensor,
    weight_shape: Shape,
    options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
    let dim_c = input.meta.num_dims() - 1;

    check_pointwise(&weight_shape[1..dim_c], &options)?;

    let input = reshape_input(input); // [(NHW), C_in] : [M, N]
    let out_grad = reshape_input(out_grad); // [(NHW), C_out] : [M, K]
    let dtype = out_grad.dtype;

    // A metadata swap, so the matmul reads the gradient K-major rather than a
    // transposed copy of it being made.
    let out_grad = swap_dims(out_grad, 0, 1); // [C_out, M]

    let grad = matmul(out_grad, input, None, MatmulStrategy::default(), dtype)?; // [C_out, C_in]

    // Only unit kernel dimensions are being reinserted, so nothing moves.
    Ok(reshape(grad, weight_shape)) // [C_out, 1, .., 1, C_in]
}

/// How few rows of the contraction a single piece may be left with.
///
/// The cut is worth making because it buys parallelism, and stops being worth
/// making once each piece is too short to amortise its own launch. Flat across
/// every piece count [`MAX_SPLIT`] leaves reachable, on the shapes measured, so
/// the exact value is not delicate.
const MIN_SPLIT_ROWS: usize = 2048;

/// The most pieces to cut into, so that the reduction putting them back stays
/// small next to the matmul that produced them.
const MAX_SPLIT: usize = 64;

/// How many pieces to cut a weight gradient's contraction into, or `None` when
/// cutting it does not apply.
///
/// A weight gradient contracts over every pixel in the batch, so `k` is enormous
/// against an output that is only `c_out` by `c_in`. A matmul kernel gives each
/// output element the whole of `k`, which leaves a device with far more lanes
/// than there are output elements mostly idle — the arithmetic is fine and the
/// *shape* is wrong. Cutting `k` into independent pieces multiplies the work
/// items by the cut and costs one small reduction to put back.
///
/// Declines a contraction too short to be worth cutting, and one that none of
/// the cuts it considers divides — the search tries powers of two only, which is
/// where `k = batch * height * width` almost always has its factors, so the rare
/// `k` whose only equal cuts are odd is left uncut rather than sent down a shape
/// nothing measured. Whether the cut *pays* is left to autotune, which measures
/// it against the uncut form on the actual shape — a guess about where the
/// crossover lies would only be a worse version of that measurement.
///
/// The count has to divide `k` exactly, since the whole point is that the
/// reshape splitting it is free.
fn split_count(k: usize) -> Option<usize> {
    if k < MIN_SPLIT_ROWS * 2 {
        return None;
    }

    let by_rows = k / MIN_SPLIT_ROWS;
    let ceiling = Ord::min(by_rows, MAX_SPLIT);

    // Powers of two downward, so the split divides `k` and the pieces are
    // equal. `k` is `batch * height * width` and usually has many factors of
    // two, but nothing guarantees it, so this can come back empty.
    (1..=ceiling.ilog2())
        .rev()
        .map(|log| 1usize << log)
        .find(|split| k.is_multiple_of(*split))
}

/// The gradient with respect to a 1x1 convolution's weight, with the
/// contraction cut into independent pieces and summed.
///
/// Identical arithmetic to [`wgrad_im2col_1x1`] up to the order the products
/// are added in, and the same single matmul underneath — only batched, over a
/// `k` that has been cut. See [`split_count`] for why that is worth doing.
///
/// Registered beside the uncut form rather than replacing it, so that autotune
/// decides per shape: the cut is a large win where the output is small and a
/// small loss where it is not, and which side a shape falls on is exactly the
/// kind of thing measuring answers better than a rule.
pub fn wgrad_im2col_1x1_split<const N: usize>(
    input: CubeTensor,
    out_grad: CubeTensor,
    weight_shape: Shape,
    options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
    let dim_c = input.meta.num_dims() - 1;

    check_pointwise(&weight_shape[1..dim_c], &options)?;

    // Every way of bowing out below ends in the uncut form rather than in an
    // `Err`, so that this candidate declines exactly what [`wgrad_im2col_1x1`]
    // declines and nothing more. What it would otherwise decline on turns on
    // `k`, and the autotune key holds the spatial dimensions only anchored: a
    // shape that declines can share a key with one that did not, and the tuner
    // unwraps whatever it already picked, so declining on a cached hit aborts
    // the process.
    let uncut = {
        let args = (
            input.clone(),
            out_grad.clone(),
            weight_shape.clone(),
            options.clone(),
        );
        move || wgrad_im2col_1x1::<N>(args.0, args.1, args.2, args.3)
    };

    let rows: usize = input.meta.shape()[..dim_c].iter().product();
    let Some(split) = split_count(rows) else {
        return uncut();
    };
    let per = rows / split;

    let input = reshape_input(input); // [M, C_in]
    let out_grad = reshape_input(out_grad); // [M, C_out]
    let dtype = out_grad.dtype;

    let in_channels = input.meta.shape()[1];
    let out_channels = out_grad.meta.shape()[1];

    // `[M, C]` -> `[split, M / split, C]`. Free: the contraction is the leading
    // axis of both operands, so cutting it only inserts a dimension.
    let input = reshape(input, Shape::new([split, per, in_channels]));
    let out_grad = reshape(out_grad, Shape::new([split, per, out_channels]));

    // `[split, C_out, M / split] @ [split, M / split, C_in]`, a stride swap on
    // the gradient as in the uncut form.
    let out_grad = swap_dims(out_grad, 1, 2);
    let Ok(partials) = matmul(out_grad, input, None, MatmulStrategy::default(), dtype) else {
        return uncut();
    };

    // `[split, C_out, C_in]` -> `[1, C_out, C_in]`. Small next to the matmul:
    // the pieces are the only thing being added, not the contraction.
    let grad = reduce_dim(
        partials,
        None,
        0,
        KernelReduceStrategy::default(),
        ReduceOperationConfig::Sum,
    );
    // The axis is in range, so only the strategy or the dtype can refuse here —
    // and the uncut form adds the same products inside its own matmul.
    let Ok(grad) = grad else {
        return uncut();
    };

    Ok(reshape(grad, weight_shape))
}

/// The input laid out as the matrix a weight gradient contracts against:
/// `[(batch, ..out spatial), (..kernel, channels)]`.
///
/// `columns[(n, ..o), (..k, c)] = input[n, ..o * stride + k * dilation -
/// padding.., c]`, and zero where that reads outside the image — a padded
/// position contributes nothing to the gradient, so zero *is* the answer.
///
/// Built from one assignment per kernel tap rather than a kernel of its own.
/// Each tap owns a contiguous block of columns and reads a sub-rectangle of the
/// image, so at unit stride its source is a plain slice — metadata only. A
/// strided convolution pays a gather per tap on top.
///
/// The column axis is ordered `(..kernel, channels)` deliberately: that is the
/// weight's own NHWC layout, so what the matmul produces needs reshaping and
/// not permuting.
///
/// **This materialises.** The column matrix holds every pixel once per tap that
/// reads it — nine times over for a 3x3 — which is what im2col costs everywhere
/// and what buys the contraction a shape a matmul can hold.
fn im2col<const N: usize>(
    input: CubeTensor,
    out_shape: &[usize],
    kernel_shape: &[usize],
    options: &ConvOptions<N>,
) -> CubeTensor {
    let rank = input.meta.num_dims();
    let dim_c = rank - 1;

    let batch = input.meta.shape()[0];
    let channels = input.meta.shape()[dim_c];
    let in_shape = input.meta.shape()[1..dim_c].to_vec();

    let taps: usize = kernel_shape.iter().product();

    let mut columns_shape = vec![batch];
    columns_shape.extend(out_shape.iter().copied());
    columns_shape.push(taps * channels);

    // Every tap is planned before anything is allocated, because whether *any*
    // of them is clipped is what decides if the column matrix has to be zeroed.
    let mut blocks = Vec::with_capacity(taps);
    let mut clipped = false;

    for tap in 0..taps {
        // The tap's index per spatial dimension, innermost varying fastest —
        // the order the column axis is laid out in.
        let mut rest = tap;
        let mut offsets = vec![0usize; N];
        for axis in (0..N).rev() {
            offsets[axis] = rest % kernel_shape[axis];
            rest /= kernel_shape[axis];
        }

        let mut source = vec![Slice::from(0..batch)];
        let mut target = vec![Slice::from(0..batch)];
        let mut covers_nothing = false;

        for axis in 0..N {
            let stride = options.stride[axis] as isize;
            // Where this tap reads for output zero. Negative under padding.
            let base = (offsets[axis] * options.dilation[axis]) as isize
                - options.padding_begin()[axis] as isize;
            let extent = in_shape[axis] as isize;

            // The outputs whose read lands inside the image. Everything else is
            // padding, and zero is already the answer there.
            // Rounded up by hand: `isize::div_ceil` is not stable, and both
            // operands are positive here.
            let first = match base >= 0 {
                true => 0,
                false => (-base + stride - 1) / stride,
            };
            let last = match extent - 1 - base {
                reach if reach < 0 => 0,
                reach => Ord::min(out_shape[axis] as isize, reach / stride + 1),
            };

            if last <= first {
                covers_nothing = true;
                break;
            }

            clipped |= first > 0 || last < out_shape[axis] as isize;

            // Exactly `last - first` elements: the end is one past the last one
            // the step actually lands on, not one past the range it spans.
            let start = first * stride + base;
            source.push(Slice {
                start,
                end: Some(start + (last - first - 1) * stride + 1),
                step: stride,
            });
            target.push(Slice {
                start: first,
                end: Some(last),
                step: 1,
            });
        }

        // A tap that reads outside the image everywhere — a kernel wider than
        // the padded image. Its columns stay zero.
        if covers_nothing {
            clipped = true;
            continue;
        }

        source.push(Slice::from(0..channels));
        target.push(Slice::from(tap * channels..(tap + 1) * channels));

        blocks.push((source, target));
    }

    // Only a clipped tap leaves a hole, and with no padding there is none: the
    // taps together write every column, so the fill would be a full pass over
    // the largest buffer here that nothing reads back.
    let mut columns = match clipped {
        true => zeros_client(
            input.client.clone(),
            input.device.clone(),
            columns_shape.into(),
            input.dtype,
        ),
        false => empty_device_dtype(
            input.client.clone(),
            input.device.clone(),
            columns_shape.into(),
            input.dtype,
        ),
    };

    for (source, target) in blocks {
        let block = slice_with_steps(input.clone(), &source);
        columns = slice_assign(columns, &target, block);
    }

    columns
}

/// The gradient with respect to a dense convolution's weight, as one matmul
/// over the input's columns.
///
/// `conv_weight_grad_no_groups` computes the same thing by convolving the input
/// *by the output gradient*, which makes the gradient the kernel — so the
/// convolution it submits has a kernel the size of the whole feature map and
/// channels numbering only the batch. That is the worst shape a direct
/// convolution kernel can be handed, and on a device whose dtype has no
/// accelerated matmul it is the only candidate that does not decline. A 3x3
/// over a 128x128 map at batch 4 takes about 20 ms that way, against 1.5 ms for
/// the *data* gradient of the same convolution.
///
/// Laid out as columns it is `[c_out, m] @ [m, ..kernel * c_in]`, a matmul —
/// and one whose contraction is enormous against a small output, so
/// [`split_count`] applies for the same reason it does to the pointwise case.
///
/// The cost is the materialisation: see [`im2col`]. That is the trade this
/// makes, and it is why it is offered to autotune rather than taken as a rule.
pub fn wgrad_im2col<const N: usize>(
    input: CubeTensor,
    out_grad: CubeTensor,
    weight_shape: Shape,
    options: ConvOptions<N>,
) -> Result<CubeTensor, ConvSetupError> {
    let rank = input.meta.num_dims();
    let dim_c = rank - 1;

    // Both declines read fields the autotune key holds exactly. A decline that
    // turned on the spatial dimensions or the batch would be unsound: the key
    // anchors those, so a shape that declines can share a key with one that did
    // not, and the tuner unwraps whatever it already picked.
    if options.groups != 1 {
        return Err(ConvSetupError::Groups(options.groups));
    }
    // A pointwise convolution's columns are a copy of the input feeding the
    // matmul `wgrad_im2col_1x1` already runs without one, so this can only lose.
    if check_pointwise(&weight_shape[1..dim_c], &options).is_ok() {
        return Err(ConvSetupError::Unknown);
    }

    let out_channels = weight_shape[0];
    let in_channels = input.meta.shape()[dim_c];
    let kernel_shape = weight_shape[1..dim_c].to_vec();
    let out_shape = out_grad.meta.shape()[1..dim_c].to_vec();

    let cols = kernel_shape.iter().product::<usize>() * in_channels;
    let rows = out_grad.meta.shape()[..dim_c].iter().product::<usize>();

    let columns = im2col::<N>(input, &out_shape, &kernel_shape, &options);
    // `[batch, ..out spatial, cols]` -> `[m, cols]`. Free: only leading
    // dimensions merge, and they are dense in that order.
    let columns = reshape(columns, Shape::new([rows, cols]));

    let out_grad = reshape_input(out_grad); // [m, c_out]
    let dtype = out_grad.dtype;

    // The uncut form, which every way of bowing out of the cut below ends in
    // rather than in an `Err`. What the cut turns on is `rows`, and the key
    // holds the spatial dimensions and the batch only anchored: a shape that
    // declines can share a key with one that did not, and the tuner unwraps
    // whatever it already picked, so declining on a cached hit aborts.
    let uncut = |columns: CubeTensor, out_grad: CubeTensor| {
        let out_grad = swap_dims(out_grad, 0, 1); // [c_out, m]
        matmul(out_grad, columns, None, MatmulStrategy::default(), dtype)
    };

    let grad = match split_count(rows) {
        // `[split, c_out, m / split] @ [split, m / split, cols]`, the partials
        // summed. Every reshape is free, as in `wgrad_im2col_1x1_split`.
        Some(split) => {
            let per = rows / split;
            let cut = reshape(columns.clone(), Shape::new([split, per, cols]));
            let grad = reshape(out_grad.clone(), Shape::new([split, per, out_channels]));
            let grad = swap_dims(grad, 1, 2);

            let partials = matmul(grad, cut, None, MatmulStrategy::default(), dtype)
                .ok()
                .and_then(|partials| {
                    reduce_dim(
                        partials,
                        None,
                        0,
                        KernelReduceStrategy::default(),
                        ReduceOperationConfig::Sum,
                    )
                    .ok()
                });

            match partials {
                Some(grad) => grad,
                None => uncut(columns, out_grad)?,
            }
        }
        None => uncut(columns, out_grad)?,
    };

    // `[c_out, ..kernel * c_in]` -> `[c_out, ..kernel, c_in]`, which is the
    // order the columns were laid out in.
    Ok(reshape(grad, weight_shape))
}