cubek_convolution/kernels/algorithm/
specialized.rs1use 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
25pub struct SpecializedConv<TMM: TileMatmulFamily, L: AsyncPartialLoadingStrategy<RuntimeArgs>> {
27 _tmm: PhantomData<TMM>,
28 _loader: PhantomData<L>,
29}
30
31pub 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}