1use cubecl::{Runtime, client::ComputeClient};
8use cubek_matmul::{
9 definition::MatmulElems,
10 multi_level::{
11 components::tile::TileMatmulKind, definition::BatchMatmulBlueprint, routines::TilingArgs,
12 },
13 routine::{BlueprintStrategy, Routine as MatmulRoutine},
14};
15
16use crate::{components::ConvolutionOperation, definition::ConvBlueprint};
17
18fn blueprint_operation(blueprint: &ConvBlueprint) -> ConvolutionOperation {
19 match blueprint {
20 ConvBlueprint::Forward(_) => ConvolutionOperation::Forward,
21 ConvBlueprint::BackwardData(_) => ConvolutionOperation::BackwardData,
22 ConvBlueprint::BackwardWeight(_) => ConvolutionOperation::BackwardWeight,
23 }
24}
25
26use crate::{
27 components::{ConvSetupError, global::args::RuntimeArgs},
28 kernels::{backward_data, backward_weight, forward},
29 launch::{
30 ConvAlgorithm, ConvolutionArgs, ConvolutionInputs, Strategy, strategy::AcceleratedTileKind,
31 },
32 routines::{
33 Routine,
34 simple::{
35 SimpleAsyncCyclicConv, SimpleAsyncStridedConv, SimpleAsyncTmaConv,
36 SimpleSyncCyclicConv, SimpleSyncStridedConv, SimpleSyncTilewiseConv,
37 },
38 specialized::{
39 SpecializedAsyncCyclicConv, SpecializedAsyncStridedConv, SpecializedTmaConv,
40 },
41 },
42};
43
44pub(crate) fn tile_kind_to_dispatch(kind: AcceleratedTileKind) -> TileMatmulKind {
46 match kind {
47 AcceleratedTileKind::Cmma => TileMatmulKind::Cmma,
48 AcceleratedTileKind::Mma => TileMatmulKind::Mma,
49 }
50}
51
52#[allow(clippy::result_large_err)]
58pub fn launch_ref<R: Runtime, const N_SPATIAL: usize>(
59 strategy: &Strategy,
60 client: &ComputeClient<R>,
61 inputs: ConvolutionInputs<R>,
62 args: ConvolutionArgs<N_SPATIAL>,
63 dtypes: MatmulElems,
64) -> Result<(), ConvSetupError> {
65 let (algorithm, tile_kind, forced_matmul) = match strategy {
66 Strategy::Inferred {
67 algorithm,
68 tile_kind,
69 } => (*algorithm, *tile_kind, None),
70 Strategy::Forced {
71 algorithm,
72 blueprint,
73 } => {
74 debug_assert_eq!(
75 inputs.operation(),
76 blueprint_operation(blueprint),
77 "Strategy::Forced blueprint variant does not match the inputs operation",
78 );
79 let matmul = blueprint.matmul().clone();
80 (*algorithm, AcceleratedTileKind::Cmma, Some(matmul))
84 }
85 };
86
87 if inputs.operation() == ConvolutionOperation::BackwardData
89 && algorithm == ConvAlgorithm::SimpleAsyncTma
90 {
91 return Err(crate::kernels::backward_data::launch::unsupported_tma_error());
92 }
93
94 dispatch_routine::<R, N_SPATIAL>(
95 algorithm,
96 tile_kind,
97 forced_matmul,
98 client,
99 inputs,
100 args,
101 dtypes,
102 )
103}
104
105#[allow(clippy::result_large_err, clippy::too_many_arguments)]
108fn dispatch_routine<R: Runtime, const N_SPATIAL: usize>(
109 algorithm: ConvAlgorithm,
110 tile_kind: AcceleratedTileKind,
111 forced_matmul: Option<BatchMatmulBlueprint>,
112 client: &ComputeClient<R>,
113 inputs: ConvolutionInputs<R>,
114 args: ConvolutionArgs<N_SPATIAL>,
115 dtypes: MatmulElems,
116) -> Result<(), ConvSetupError> {
117 let kind = tile_kind_to_dispatch(tile_kind);
118 match algorithm {
119 ConvAlgorithm::SimpleSyncCyclic => dispatch_inputs::<R, N_SPATIAL, SimpleSyncCyclicConv>(
120 client,
121 inputs,
122 args,
123 kind,
124 forced_matmul,
125 dtypes,
126 ),
127 ConvAlgorithm::SimpleSyncStrided => dispatch_inputs::<R, N_SPATIAL, SimpleSyncStridedConv>(
128 client,
129 inputs,
130 args,
131 kind,
132 forced_matmul,
133 dtypes,
134 ),
135 ConvAlgorithm::SimpleSyncTilewise => {
136 dispatch_inputs::<R, N_SPATIAL, SimpleSyncTilewiseConv>(
137 client,
138 inputs,
139 args,
140 kind,
141 forced_matmul,
142 dtypes,
143 )
144 }
145 ConvAlgorithm::SimpleAsyncCyclic => dispatch_inputs::<R, N_SPATIAL, SimpleAsyncCyclicConv>(
146 client,
147 inputs,
148 args,
149 kind,
150 forced_matmul,
151 dtypes,
152 ),
153 ConvAlgorithm::SimpleAsyncStrided => {
154 dispatch_inputs::<R, N_SPATIAL, SimpleAsyncStridedConv>(
155 client,
156 inputs,
157 args,
158 kind,
159 forced_matmul,
160 dtypes,
161 )
162 }
163 ConvAlgorithm::SimpleAsyncTma => dispatch_inputs::<R, N_SPATIAL, SimpleAsyncTmaConv>(
164 client,
165 inputs,
166 args,
167 kind,
168 forced_matmul,
169 dtypes,
170 ),
171 ConvAlgorithm::SpecializedAsyncCyclic => {
172 dispatch_inputs::<R, N_SPATIAL, SpecializedAsyncCyclicConv>(
173 client,
174 inputs,
175 args,
176 kind,
177 forced_matmul,
178 dtypes,
179 )
180 }
181 ConvAlgorithm::SpecializedAsyncStrided => {
182 dispatch_inputs::<R, N_SPATIAL, SpecializedAsyncStridedConv>(
183 client,
184 inputs,
185 args,
186 kind,
187 forced_matmul,
188 dtypes,
189 )
190 }
191 ConvAlgorithm::SpecializedTma => dispatch_inputs::<R, N_SPATIAL, SpecializedTmaConv>(
192 client,
193 inputs,
194 args,
195 kind,
196 forced_matmul,
197 dtypes,
198 ),
199 }
200}
201
202#[allow(clippy::result_large_err, clippy::too_many_arguments)]
208fn dispatch_inputs<
209 R: Runtime,
210 const N_SPATIAL: usize,
211 Rt: Routine<Blueprint = BatchMatmulBlueprint>,
212>(
213 client: &ComputeClient<R>,
214 inputs: ConvolutionInputs<R>,
215 args: ConvolutionArgs<N_SPATIAL>,
216 tile_matmul: TileMatmulKind,
217 forced_matmul: Option<BatchMatmulBlueprint>,
218 dtypes: MatmulElems,
219) -> Result<(), ConvSetupError>
220where
221 Rt::Args: forward::args::ConcreteArgs<Rt::MatmulRoutine>
222 + backward_data::args::ConcreteArgs<Rt::MatmulRoutine>
223 + backward_weight::args::ConcreteArgs<Rt::MatmulRoutine>,
224 Rt::Strategy: TilingArgs,
225{
226 let blueprint_strategy = build_blueprint_strategy::<Rt>(tile_matmul, forced_matmul);
227
228 match inputs {
229 ConvolutionInputs::Forward {
230 input,
231 weight,
232 bias,
233 out,
234 } => forward::launch::launch_internal::<R, N_SPATIAL, Rt>(
235 client,
236 input,
237 weight,
238 bias,
239 out,
240 args,
241 &blueprint_strategy,
242 dtypes,
243 ),
244 ConvolutionInputs::BackwardData {
245 out_grad,
246 weights,
247 in_grad,
248 } => backward_data::launch::launch_internal::<R, N_SPATIAL, Rt>(
249 client,
250 out_grad,
251 weights,
252 in_grad,
253 args,
254 &blueprint_strategy,
255 dtypes,
256 ),
257 ConvolutionInputs::BackwardWeight {
258 input,
259 out_grad,
260 weight_grad,
261 } => backward_weight::launch::launch_internal::<R, N_SPATIAL, Rt>(
262 client,
263 input,
264 out_grad,
265 weight_grad,
266 args,
267 &blueprint_strategy,
268 dtypes,
269 ),
270 }
271}
272
273fn build_blueprint_strategy<Rt: Routine<Blueprint = BatchMatmulBlueprint>>(
277 tile_matmul: TileMatmulKind,
278 forced_matmul: Option<BatchMatmulBlueprint>,
279) -> BlueprintStrategy<RuntimeArgs, Rt::MatmulRoutine>
280where
281 Rt::Strategy: TilingArgs,
282{
283 match forced_matmul {
284 Some(matmul) => BlueprintStrategy::Forced(matmul),
285 None => {
286 let mut s = <Rt::MatmulRoutine as MatmulRoutine<RuntimeArgs>>::Strategy::default();
287 s.set_tile_matmul(tile_matmul);
288 BlueprintStrategy::Inferred(s)
289 }
290 }
291}