cubek_convolution/kernels/forward/
launch.rs1use 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#[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 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 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}