Skip to main content

cubek_convolution/kernels/algorithm/
specialized.rs

1use std::marker::PhantomData;
2
3use cubecl::{
4    Runtime, client::ComputeClient, ir::StorageType, prelude::TensorHandleRef, server::LaunchError,
5    std::tensor::TensorHandle,
6};
7use cubek_matmul::{
8    components::{
9        global::read::{AsyncPartialLoadingStrategy, async_partial_tma::AsyncPartialTmaLoading},
10        tile::{TileMatmulFamily, io::Strided},
11    },
12    definition::AvailableLineSizes,
13    launch::{TensorArgs, TensorMapArgs},
14    routines::specialized::SpecializedAlgorithm,
15};
16
17use crate::{
18    algorithm::{Algorithm, into_tensor_handle, into_tensor_handle_tma},
19    components::{
20        ConvolutionOperation,
21        global::{args::RuntimeArgs, read::strategy::sync_bias::SyncBiasLoading},
22    },
23};
24
25/// Cmma convolution
26pub struct SpecializedConv<TMM: TileMatmulFamily, L: AsyncPartialLoadingStrategy<RuntimeArgs>> {
27    _tmm: PhantomData<TMM>,
28    _loader: PhantomData<L>,
29}
30
31// pub type SpecializedCyclicConv<TMM> =
32//     SpecializedConv<TMM, AsyncPartialCyclicLoading<ColMajorTilingOrder>>;
33// pub type SpecializedStridedConv<TMM> = SpecializedConv<TMM, AsyncPartialStridedLoading>;
34
35pub struct SpecializedTmaConv<TMM: TileMatmulFamily> {
36    _tmm: PhantomData<TMM>,
37}
38
39impl<
40    TMM: TileMatmulFamily<
41            LhsTile = Strided,
42            RhsTile = Strided,
43            AccTile = Option<Strided>,
44            OutTile = Strided,
45        >,
46    L: AsyncPartialLoadingStrategy<RuntimeArgs, TileKind = Strided>,
47> Algorithm for SpecializedConv<TMM, L>
48{
49    type Routine = SpecializedAlgorithm<TMM, L, SyncBiasLoading>;
50    type Args = TensorArgs<RuntimeArgs>;
51    const IS_SPECIALIZED: bool = true;
52
53    fn into_tensor_handle<R: Runtime>(
54        client: &ComputeClient<R>,
55        handle: &TensorHandleRef<'_, R>,
56        dtype: StorageType,
57        _operation: ConvolutionOperation,
58    ) -> Result<TensorHandle<R>, LaunchError> {
59        into_tensor_handle(client, handle, dtype)
60    }
61}
62
63impl<
64    TMM: TileMatmulFamily<
65            LhsTile = Strided,
66            RhsTile = Strided,
67            AccTile = Option<Strided>,
68            OutTile = Strided,
69        >,
70> Algorithm for SpecializedTmaConv<TMM>
71{
72    type Routine = SpecializedAlgorithm<TMM, AsyncPartialTmaLoading, SyncBiasLoading>;
73    type Args = TensorMapArgs<RuntimeArgs>;
74    const IS_SPECIALIZED: bool = true;
75
76    fn into_tensor_handle<R: Runtime>(
77        client: &ComputeClient<R>,
78        handle: &TensorHandleRef<'_, R>,
79        dtype: StorageType,
80        operation: ConvolutionOperation,
81    ) -> Result<TensorHandle<R>, LaunchError> {
82        into_tensor_handle_tma(client, handle, dtype, operation)
83    }
84
85    fn filter_line_sizes(line_sizes: AvailableLineSizes) -> AvailableLineSizes {
86        AvailableLineSizes {
87            lhs: vec![1],
88            rhs: vec![1],
89            out: line_sizes.out,
90        }
91    }
92}