cubek_convolution/routines/
simple.rs1use 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
35pub 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}