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