Skip to main content

cubek_convolution/routines/
simple.rs

1use cubecl::{
2    server::LaunchError,
3    {Runtime, client::ComputeClient, ir::StorageType, prelude::TensorBinding},
4};
5use cubek_matmul::components::global::read::FullLoadingStrategy;
6use cubek_matmul::components::global::read::sync_full_cyclic::SyncFullCyclicLoading;
7use cubek_matmul::{
8    args::{TensorArgs, TensorMapArgs},
9    definition::{AvailableVectorSizes, BatchMatmulBlueprint},
10};
11use cubek_matmul::{
12    components::global::read::{
13        async_full_tma::AsyncFullTmaLoading, sync_full_strided::SyncFullStridedLoading,
14        sync_full_tilewise::SyncFullTilewiseLoading,
15    },
16    routines::batch::simple::{SimpleAlgorithm, SimpleArgs},
17};
18use cubek_std::tile::{ColMajorTilingOrder, RowMajorTilingOrder};
19use std::marker::PhantomData;
20
21use crate::{
22    components::{
23        ConvolutionOperation,
24        global::{
25            args::RuntimeArgs,
26            read::strategy::{
27                async_full_cyclic::AsyncFullCyclicLoading,
28                async_full_strided::AsyncFullStridedLoading, sync_bias::SyncBiasLoading,
29            },
30        },
31    },
32    routines::{Routine, contiguous_pitched_layout, into_tensor_handle_tma},
33};
34
35/// Cmma convolution
36pub struct SimpleConv<LL: FullLoadingStrategy<RuntimeArgs>, LR: FullLoadingStrategy<RuntimeArgs>> {
37    _loader: PhantomData<(LL, LR)>,
38}
39
40pub type SimpleSyncCyclicConv = SimpleConv<
41    SyncFullCyclicLoading<RowMajorTilingOrder>,
42    SyncFullCyclicLoading<ColMajorTilingOrder>,
43>;
44pub type SimpleSyncStridedConv = SimpleConv<SyncFullStridedLoading, SyncFullStridedLoading>;
45pub type SimpleSyncTilewiseConv = SimpleConv<
46    SyncFullTilewiseLoading<RowMajorTilingOrder>,
47    SyncFullTilewiseLoading<ColMajorTilingOrder>,
48>;
49pub type SimpleAsyncCyclicConv = SimpleConv<
50    AsyncFullCyclicLoading<RowMajorTilingOrder>,
51    AsyncFullCyclicLoading<ColMajorTilingOrder>,
52>;
53pub type SimpleAsyncStridedConv = SimpleConv<AsyncFullStridedLoading, AsyncFullStridedLoading>;
54
55pub struct SimpleAsyncTmaConv;
56
57impl<
58    LL: FullLoadingStrategy<RuntimeArgs>,
59    LR: FullLoadingStrategy<RuntimeArgs, SyncStrategy = LL::SyncStrategy>,
60> Routine for SimpleConv<LL, LR>
61{
62    type Blueprint = BatchMatmulBlueprint;
63    type Strategy = SimpleArgs;
64    type MatmulRoutine = SimpleAlgorithm<LL, LR, SyncBiasLoading>;
65    type Args = TensorArgs<RuntimeArgs>;
66
67    fn correct_layout<R: Runtime>(
68        client: &ComputeClient<R>,
69        handle: TensorBinding<R>,
70        dtype: StorageType,
71        _operation: ConvolutionOperation,
72    ) -> Result<TensorBinding<R>, LaunchError> {
73        contiguous_pitched_layout(client, handle, dtype)
74    }
75}
76
77impl Routine for SimpleAsyncTmaConv {
78    type Blueprint = BatchMatmulBlueprint;
79    type Strategy = SimpleArgs;
80    type MatmulRoutine = SimpleAlgorithm<AsyncFullTmaLoading, AsyncFullTmaLoading, SyncBiasLoading>;
81    type Args = TensorMapArgs<RuntimeArgs>;
82
83    fn correct_layout<R: Runtime>(
84        client: &ComputeClient<R>,
85        handle: TensorBinding<R>,
86        dtype: StorageType,
87        operation: ConvolutionOperation,
88    ) -> Result<TensorBinding<R>, LaunchError> {
89        into_tensor_handle_tma(client, handle, dtype, operation)
90    }
91
92    fn filter_vector_sizes(vector_sizes: AvailableVectorSizes) -> AvailableVectorSizes {
93        AvailableVectorSizes {
94            lhs: vec![1],
95            rhs: vec![1],
96            out: vector_sizes.out,
97        }
98    }
99}