Skip to main content

cubek_convolution/kernels/forward/
launch.rs

1use crate::{
2    components::{
3        ConvSetupError, ConvolutionOperation, ConvolutionProblem, Dimensionality,
4        global::args::RuntimeArgs,
5    },
6    forward::args::ConcreteArgs,
7    kernels::forward::selector::launch_kernel_concrete,
8    launch::ConvolutionArgs,
9    routines::Routine,
10};
11use cubecl::{Runtime, client::ComputeClient, prelude::*};
12use cubek_matmul::{
13    definition::{AvailableVectorSizes, MatmulElems},
14    routine::BlueprintStrategy,
15};
16use cubek_std::{InputBinding, MatrixLayout};
17
18/// Forward-convolution dispatch helper.
19///
20/// Called by `cubek_convolution::launch_ref` after the routine and
21/// blueprint-strategy have been resolved. Not meant for direct external use.
22#[allow(clippy::result_large_err, clippy::too_many_arguments)]
23pub(crate) fn launch_internal<R: Runtime, const N_SPATIAL: usize, Rt: Routine>(
24    client: &ComputeClient<R>,
25    input: InputBinding<R>,
26    weight: InputBinding<R>,
27    bias: Option<InputBinding<R>>,
28    out: TensorBinding<R>,
29    args: ConvolutionArgs<N_SPATIAL>,
30    blueprint_strategy: &BlueprintStrategy<RuntimeArgs, Rt::MatmulRoutine>,
31    dtypes: MatmulElems,
32) -> Result<(), ConvSetupError>
33where
34    Rt::Args: ConcreteArgs<Rt::MatmulRoutine>,
35{
36    let ConvolutionArgs {
37        stride,
38        padding,
39        dilation,
40    } = args;
41
42    let dimensionality = match N_SPATIAL {
43        1 => Dimensionality::Dim1,
44        2 => Dimensionality::Dim2,
45        3 => Dimensionality::Dim3,
46        other => unimplemented!("Unsupported dimensionality {other}"),
47    };
48
49    launch_with_routine::<R, Rt>(
50        client,
51        input,
52        weight,
53        bias,
54        out,
55        (&stride, &padding, &dilation),
56        dimensionality,
57        blueprint_strategy,
58        dtypes,
59    )
60}
61
62#[allow(clippy::too_many_arguments)]
63fn launch_with_routine<R: Runtime, Rt: Routine>(
64    client: &ComputeClient<R>,
65    input: InputBinding<R>,
66    weight: InputBinding<R>,
67    bias: Option<InputBinding<R>>,
68    out: TensorBinding<R>,
69    (stride, padding, dilation): (&[usize], &[usize], &[usize]),
70    dimensionality: Dimensionality,
71    blueprint_strategy: &BlueprintStrategy<RuntimeArgs, Rt::MatmulRoutine>,
72    dtypes: MatmulElems,
73) -> Result<(), ConvSetupError>
74where
75    Rt::Args: ConcreteArgs<Rt::MatmulRoutine>,
76{
77    let rank = input.data().shape.len();
78    let dim_c = rank - 1;
79
80    let n = input.data().shape[0];
81    let c = input.data().shape[dim_c];
82
83    let out_c = weight.data().shape[0];
84
85    let in_shape = &input.data().shape[1..dim_c];
86    let kernel_shape = &weight.data().shape[1..dim_c];
87    let out_shape = &out.shape[1..dim_c];
88
89    let op = ConvolutionOperation::Forward;
90
91    let input_data = Rt::correct_layout(client, input.clone().into_data(), dtypes.lhs_global, op)?;
92    let weight_data =
93        Rt::correct_layout(client, weight.clone().into_data(), dtypes.rhs_global, op)?;
94
95    let mut input = input.clone();
96    let mut weight = weight.clone();
97
98    *input.data_mut() = input_data;
99    *weight.data_mut() = weight_data;
100
101    let address_type = input
102        .required_address_type()
103        .max(weight.required_address_type())
104        .max(
105            bias.clone()
106                .map(|bias| bias.required_address_type())
107                .unwrap_or_default(),
108        )
109        .max(out.required_address_type(dtypes.acc_global.size()));
110
111    let problem = ConvolutionProblem {
112        m: n * out_shape.iter().product::<usize>(),
113        n: out_c,
114        k: c * kernel_shape.iter().product::<usize>(),
115        lhs_strides: input.data().strides.clone(),
116        rhs_strides: weight.data().strides.clone(),
117        lhs_layout: MatrixLayout::RowMajor,
118        rhs_layout: MatrixLayout::ColMajor,
119        kernel_size: kernel_shape.iter().map(|it| *it as u32).collect(),
120        stride: stride.iter().map(|it| *it as u32).collect(),
121        padding: padding.iter().map(|it| *it as i32).collect(),
122        dilation: dilation.iter().map(|it| *it as u32).collect(),
123
124        batches: n,
125        in_shape: in_shape.into(),
126        out_shape: out_shape.into(),
127        channels: c,
128        out_channels: out_c,
129
130        padded_channels: c,
131        operation: op,
132
133        dimensionality,
134        global_dtypes: dtypes.as_global_elems(),
135        address_type,
136    };
137
138    launch_kernel::<R, Rt>(
139        client,
140        input,
141        weight,
142        bias,
143        out,
144        problem,
145        blueprint_strategy,
146        dtypes,
147    )
148}
149
150#[allow(clippy::result_large_err, clippy::too_many_arguments)]
151pub fn launch_kernel<R: Runtime, Rt: Routine>(
152    client: &ComputeClient<R>,
153    input: InputBinding<R>,
154    weight: InputBinding<R>,
155    bias: Option<InputBinding<R>>,
156    out: TensorBinding<R>,
157    problem: ConvolutionProblem,
158    blueprint_strategy: &BlueprintStrategy<RuntimeArgs, Rt::MatmulRoutine>,
159    dtypes: MatmulElems,
160) -> Result<(), ConvSetupError>
161where
162    Rt::Args: ConcreteArgs<Rt::MatmulRoutine>,
163{
164    // Shape/strides are treated as k-major, with the last dim always being the contiguous one.
165    // So for the sake of selecting a vector size, the shape/strides are always row-major.
166    let vector_sizes = AvailableVectorSizes::from_type_sizes(
167        client,
168        input.data_elem_size(),
169        weight.data_elem_size(),
170        dtypes.acc_global.size(),
171    )
172    .filter_lhs_with_tensor(
173        &input.data().strides,
174        &input.data().shape,
175        MatrixLayout::RowMajor,
176    )
177    .filter_rhs_with_tensor(
178        &weight.data().strides,
179        &weight.data().shape,
180        MatrixLayout::RowMajor,
181    )
182    .filter_out_with_tensor(&out.strides, &out.shape);
183
184    let mut vector_sizes = Rt::filter_vector_sizes(vector_sizes).pick_max()?;
185
186    // The large vector size resulting from dequantizing ends up slower due to restrictions on
187    // algorithms. Use this as a quick and dirty fix.
188    if input.scale().is_some() {
189        vector_sizes.lhs = 1;
190    }
191    if weight.scale().is_some() {
192        vector_sizes.rhs = 1;
193    }
194
195    launch_kernel_concrete::<R, Rt::Args, Rt::MatmulRoutine>(
196        client,
197        input,
198        weight,
199        bias,
200        out,
201        problem,
202        vector_sizes,
203        blueprint_strategy,
204        &dtypes,
205    )
206}