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::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
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<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            // 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::<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<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/// Branch on operation and forward to the per-op launcher.
199///
200/// All three per-op `ConcreteArgs` traits share the same name and the same
201/// blanket impls on `TensorArgs<RuntimeArgs>` / `TensorMapArgs<RuntimeArgs>`,
202/// so the where clause simply requires an impl per operation.
203#[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
265/// Build a matmul `BlueprintStrategy` from either a forced `BatchMatmulBlueprint`
266/// (extracted from `ConvBlueprint`) or an `Inferred` strategy stamped with the
267/// requested tile-matmul kind.
268fn 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}