Skip to main content

cubek_convolution/launch/
base.rs

1//! Unified `launch_ref` entry point for the convolution kernel family.
2//!
3//! Picks the right `Routine` impl from `ConvAlgorithm`, threads the
4//! `Strategy` (`Specific` / `Forced`) into a matmul `BlueprintStrategy`, and
5//! dispatches to the per-operation helper based on `ConvolutionInputs`.
6
7use 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
44/// Map `AcceleratedTileKind` → matmul's `TileMatmulKind`.
45pub(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/// The single public convolution entry point.
53///
54/// Routes the `inputs` (whose discriminant is the operation) and `strategy`
55/// (algorithm + tile-matmul kind, optionally a forced blueprint) into the right
56/// generic `Routine` and per-operation launch helper.
57#[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            // For Forced, tile_kind is encoded inside the matmul blueprint, so
81            // the explicit tile_kind here is unused; we pass Cmma as a benign
82            // default (it gets overwritten by the forced blueprint).
83            (*algorithm, AcceleratedTileKind::Cmma, Some(matmul))
84        }
85    };
86
87    // Backward-data does not currently support the TMA reading strategy.
88    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/// Dispatch on `ConvAlgorithm` to instantiate the right concrete `Routine`
106/// generic, then forward to the per-operation helper.
107#[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/// Branch on operation and forward to the per-op launcher.
203///
204/// All three per-op `ConcreteArgs` traits share the same name and the same
205/// blanket impls on `TensorArgs<RuntimeArgs>` / `TensorMapArgs<RuntimeArgs>`,
206/// so the where clause simply requires an impl per operation.
207#[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
273/// Build a matmul `BlueprintStrategy` from either a forced `BatchMatmulBlueprint`
274/// (extracted from `ConvBlueprint`) or an `Inferred` strategy stamped with the
275/// requested tile-matmul kind.
276fn 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}