1use cubecl::client::Client;
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<const N_SPATIAL: usize>(
59 strategy: &Strategy,
60 client: &Client,
61 inputs: ConvolutionInputs,
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::<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<const N_SPATIAL: usize>(
109 algorithm: ConvAlgorithm,
110 tile_kind: AcceleratedTileKind,
111 forced_matmul: Option<BatchMatmulBlueprint>,
112 client: &Client,
113 inputs: ConvolutionInputs,
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::<N_SPATIAL, SimpleSyncCyclicConv>(
120 client,
121 inputs,
122 args,
123 kind,
124 forced_matmul,
125 dtypes,
126 ),
127 ConvAlgorithm::SimpleSyncStrided => dispatch_inputs::<N_SPATIAL, SimpleSyncStridedConv>(
128 client,
129 inputs,
130 args,
131 kind,
132 forced_matmul,
133 dtypes,
134 ),
135 ConvAlgorithm::SimpleSyncTilewise => dispatch_inputs::<N_SPATIAL, SimpleSyncTilewiseConv>(
136 client,
137 inputs,
138 args,
139 kind,
140 forced_matmul,
141 dtypes,
142 ),
143 ConvAlgorithm::SimpleAsyncCyclic => dispatch_inputs::<N_SPATIAL, SimpleAsyncCyclicConv>(
144 client,
145 inputs,
146 args,
147 kind,
148 forced_matmul,
149 dtypes,
150 ),
151 ConvAlgorithm::SimpleAsyncStrided => dispatch_inputs::<N_SPATIAL, SimpleAsyncStridedConv>(
152 client,
153 inputs,
154 args,
155 kind,
156 forced_matmul,
157 dtypes,
158 ),
159 ConvAlgorithm::SimpleAsyncTma => dispatch_inputs::<N_SPATIAL, SimpleAsyncTmaConv>(
160 client,
161 inputs,
162 args,
163 kind,
164 forced_matmul,
165 dtypes,
166 ),
167 ConvAlgorithm::SpecializedAsyncCyclic => {
168 dispatch_inputs::<N_SPATIAL, SpecializedAsyncCyclicConv>(
169 client,
170 inputs,
171 args,
172 kind,
173 forced_matmul,
174 dtypes,
175 )
176 }
177 ConvAlgorithm::SpecializedAsyncStrided => {
178 dispatch_inputs::<N_SPATIAL, SpecializedAsyncStridedConv>(
179 client,
180 inputs,
181 args,
182 kind,
183 forced_matmul,
184 dtypes,
185 )
186 }
187 ConvAlgorithm::SpecializedTma => dispatch_inputs::<N_SPATIAL, SpecializedTmaConv>(
188 client,
189 inputs,
190 args,
191 kind,
192 forced_matmul,
193 dtypes,
194 ),
195 }
196}
197
198#[allow(clippy::result_large_err, clippy::too_many_arguments)]
204fn dispatch_inputs<const N_SPATIAL: usize, Rt: Routine<Blueprint = BatchMatmulBlueprint>>(
205 client: &Client,
206 inputs: ConvolutionInputs,
207 args: ConvolutionArgs<N_SPATIAL>,
208 tile_matmul: TileMatmulKind,
209 forced_matmul: Option<BatchMatmulBlueprint>,
210 dtypes: MatmulElems,
211) -> Result<(), ConvSetupError>
212where
213 Rt::Args: forward::args::ConcreteArgs<Rt::MatmulRoutine>
214 + backward_data::args::ConcreteArgs<Rt::MatmulRoutine>
215 + backward_weight::args::ConcreteArgs<Rt::MatmulRoutine>,
216 Rt::Strategy: TilingArgs,
217{
218 let blueprint_strategy = build_blueprint_strategy::<Rt>(tile_matmul, forced_matmul);
219
220 match inputs {
221 ConvolutionInputs::Forward {
222 input,
223 weight,
224 bias,
225 out,
226 } => forward::launch::launch_internal::<N_SPATIAL, Rt>(
227 client,
228 input,
229 weight,
230 bias,
231 out,
232 args,
233 &blueprint_strategy,
234 dtypes,
235 ),
236 ConvolutionInputs::BackwardData {
237 out_grad,
238 weights,
239 in_grad,
240 } => backward_data::launch::launch_internal::<N_SPATIAL, Rt>(
241 client,
242 out_grad,
243 weights,
244 in_grad,
245 args,
246 &blueprint_strategy,
247 dtypes,
248 ),
249 ConvolutionInputs::BackwardWeight {
250 input,
251 out_grad,
252 weight_grad,
253 } => backward_weight::launch::launch_internal::<N_SPATIAL, Rt>(
254 client,
255 input,
256 out_grad,
257 weight_grad,
258 args,
259 &blueprint_strategy,
260 dtypes,
261 ),
262 }
263}
264
265fn build_blueprint_strategy<Rt: Routine<Blueprint = BatchMatmulBlueprint>>(
269 tile_matmul: TileMatmulKind,
270 forced_matmul: Option<BatchMatmulBlueprint>,
271) -> BlueprintStrategy<RuntimeArgs, Rt::MatmulRoutine>
272where
273 Rt::Strategy: TilingArgs,
274{
275 match forced_matmul {
276 Some(matmul) => BlueprintStrategy::Forced(matmul),
277 None => {
278 let mut s = <Rt::MatmulRoutine as MatmulRoutine<RuntimeArgs>>::Strategy::default();
279 s.set_tile_matmul(tile_matmul);
280 BlueprintStrategy::Inferred(s)
281 }
282 }
283}