Skip to main content

cubek_convolution/launch/
inputs.rs

1use cubecl::prelude::TensorBinding;
2use cubek_std::InputBinding;
3
4use crate::components::ConvolutionOperation;
5
6/// Spatial convolution arguments (stride / beginning padding / dilation per spatial dim).
7///
8/// End padding is represented by the spatial extent of the output binding.
9#[derive(Clone, Debug)]
10pub struct ConvolutionArgs<const N_SPATIAL: usize> {
11    pub stride: [usize; N_SPATIAL],
12    /// Padding at the beginning of each spatial dimension.
13    pub padding: [usize; N_SPATIAL],
14    pub dilation: [usize; N_SPATIAL],
15}
16
17#[allow(clippy::large_enum_variant)]
18/// Per-operation tensor bindings supplied to `launch_ref`.
19///
20/// Each variant carries exactly the bindings the corresponding operation needs.
21/// The discriminant maps 1:1 to `ConvolutionOperation`.
22pub enum ConvolutionInputs {
23    Forward {
24        input: InputBinding,
25        weight: InputBinding,
26        bias: Option<InputBinding>,
27        out: TensorBinding,
28    },
29    BackwardData {
30        out_grad: InputBinding,
31        weights: InputBinding,
32        in_grad: TensorBinding,
33    },
34    BackwardWeight {
35        input: InputBinding,
36        out_grad: InputBinding,
37        weight_grad: TensorBinding,
38    },
39}
40
41impl ConvolutionInputs {
42    pub fn operation(&self) -> ConvolutionOperation {
43        match self {
44            ConvolutionInputs::Forward { .. } => ConvolutionOperation::Forward,
45            ConvolutionInputs::BackwardData { .. } => ConvolutionOperation::BackwardData,
46            ConvolutionInputs::BackwardWeight { .. } => ConvolutionOperation::BackwardWeight,
47        }
48    }
49}