Skip to main content

cubek_std/tile/variants/instruction/
interleaved.rs

1use cubecl::prelude::*;
2
3use crate::{
4    MatrixLayout, StageIdent, TileSize,
5    tile::{SharedTile, Tile, TileKind, TileKindExpand, TileScope},
6};
7
8/// Interleaved-on-k tile. Holds just the minimal comptime data the body uses
9/// (`tile_size` and `plane_dim`); the matmul-level configuration (and any
10/// metadata not consumed by the tile body, like swizzle modes) lives in
11/// cubek-matmul as `InterleavedMatmul`.
12#[derive(CubeType)]
13pub struct InterleavedTile<N: Numeric> {
14    pub data: Array<N>,
15    #[cube(comptime)]
16    pub matrix_layout: MatrixLayout,
17    #[cube(comptime)]
18    pub tile_size: TileSize,
19    #[cube(comptime)]
20    pub plane_dim: u32,
21}
22
23#[cube]
24pub fn interleaved_allocate_lhs<L: Numeric, Sc: TileScope>(
25    #[comptime] layout: MatrixLayout,
26    #[comptime] tile_size: TileSize,
27    #[comptime] plane_dim: u32,
28) -> Tile<L, Sc> {
29    let m = tile_size.m();
30    let k = tile_size.k();
31    Tile::from_kind(TileKind::new_Interleaved(InterleavedTile::<L> {
32        data: Array::new((m * (k / plane_dim)) as usize),
33        matrix_layout: layout,
34        tile_size,
35        plane_dim,
36    }))
37}
38
39#[cube]
40pub fn interleaved_allocate_rhs<R: Numeric, Sc: TileScope>(
41    #[comptime] layout: MatrixLayout,
42    #[comptime] tile_size: TileSize,
43    #[comptime] plane_dim: u32,
44) -> Tile<R, Sc> {
45    let n = tile_size.n();
46    let k = tile_size.k();
47    Tile::from_kind(TileKind::new_Interleaved(InterleavedTile::<R> {
48        data: Array::new(((k / plane_dim) * n) as usize),
49        matrix_layout: layout,
50        tile_size,
51        plane_dim,
52    }))
53}
54
55#[cube]
56pub fn interleaved_allocate_acc<A: Numeric, Sc: TileScope>(
57    #[comptime] layout: MatrixLayout,
58    #[comptime] tile_size: TileSize,
59    #[comptime] plane_dim: u32,
60) -> Tile<A, Sc> {
61    let m = tile_size.m();
62    let n = tile_size.n();
63    Tile::from_kind(TileKind::new_Interleaved(InterleavedTile::<A> {
64        data: Array::new((m * n) as usize),
65        matrix_layout: layout,
66        tile_size,
67        plane_dim,
68    }))
69}
70
71#[cube]
72impl<A: Numeric> InterleavedTile<A> {
73    /// Executes `lhs ยท rhs`, accumulating into `self` via the plane-
74    /// interleaved-on-k matmul.
75    pub fn mma<L: Numeric, R: Numeric>(
76        &mut self,
77        lhs: &InterleavedTile<L>,
78        rhs: &InterleavedTile<R>,
79    ) {
80        interleaved_execute(
81            &lhs.data,
82            lhs.matrix_layout,
83            &rhs.data,
84            rhs.matrix_layout,
85            &mut self.data,
86            self.matrix_layout,
87            self.tile_size,
88            self.plane_dim,
89        );
90    }
91}
92
93#[cube]
94impl<N: Numeric> InterleavedTile<N> {
95    /// Copies into the interleaved tile from `source`. Supported sources:
96    /// `Shared` and `None` (zero-init).
97    pub fn copy_from<SE: Numeric, SS: Size, Sc: TileScope>(
98        &mut self,
99        source: &Tile<SE, Sc>,
100        #[comptime] ident: StageIdent,
101    ) {
102        match &source.kind {
103            TileKind::SharedTile(shared) => {
104                interleaved_load_from_shared::<SE, SS, N>(
105                    shared,
106                    &mut self.data,
107                    self.tile_size,
108                    self.plane_dim,
109                    ident,
110                );
111            }
112            TileKind::None => {
113                interleaved_load_zeros::<N>(&mut self.data, self.tile_size);
114            }
115            TileKind::Cmma(_)
116            | TileKind::Mma(_)
117            | TileKind::Register(_)
118            | TileKind::PlaneVec(_)
119            | TileKind::Interleaved(_)
120            | TileKind::Unit(_)
121            | TileKind::WhiteboxFragment(_)
122            | TileKind::RowWise(_)
123            | TileKind::Bounce(_)
124            | TileKind::Stage(_)
125            | TileKind::Partition(_)
126            | TileKind::Pipelined(_) => {
127                panic!("InterleavedTile::copy_from: unsupported source variant")
128            }
129        }
130    }
131
132    pub fn init_zero(&mut self) {
133        interleaved_load_zeros::<N>(&mut self.data, self.tile_size);
134    }
135}
136
137// ===========================================================================
138// Compute: matmul / load / write / zero-init
139// ===========================================================================
140
141#[cube]
142#[allow(clippy::too_many_arguments)]
143pub fn interleaved_execute<L: Numeric, R: Numeric, A: Numeric>(
144    lhs: &Array<L>,
145    #[comptime] lhs_layout: MatrixLayout,
146    rhs: &Array<R>,
147    #[comptime] rhs_layout: MatrixLayout,
148    acc: &mut Array<A>,
149    #[comptime] _acc_layout: MatrixLayout,
150    #[comptime] tile_size: TileSize,
151    #[comptime] plane_dim: u32,
152) {
153    let m = tile_size.m() as usize;
154    let n = tile_size.n() as usize;
155    let k = tile_size.k() as usize;
156    let plane_dim = plane_dim as usize;
157    let local_k = k / plane_dim;
158
159    let (lhs_row_count, lhs_col_count) = (m, local_k);
160    let (rhs_row_count, rhs_col_count) = (local_k, n);
161
162    #[unroll]
163    for m_ in 0..m {
164        #[unroll]
165        for n_ in 0..n {
166            #[unroll]
167            for k_ in 0..local_k {
168                let lhs_elem = A::cast_from(match lhs_layout {
169                    MatrixLayout::RowMajor => lhs[m_ * lhs_col_count + k_],
170                    MatrixLayout::ColMajor => lhs[k_ * lhs_row_count + m_],
171                });
172                let rhs_elem = A::cast_from(match rhs_layout {
173                    MatrixLayout::RowMajor => rhs[k_ * rhs_col_count + n_],
174                    MatrixLayout::ColMajor => rhs[n_ * rhs_row_count + k_],
175                });
176                acc[m_ * n + n_] += lhs_elem * rhs_elem;
177            }
178        }
179    }
180}
181
182#[cube]
183pub fn interleaved_load_from_shared<E: Numeric, ES: Size, N: Numeric>(
184    shared: &SharedTile<E>,
185    arr: &mut Array<N>,
186    #[comptime] tile_size: TileSize,
187    #[comptime] plane_dim: u32,
188    #[comptime] ident: StageIdent,
189) {
190    let shared = shared.view::<ES>();
191    let shared = &shared;
192    match ident {
193        StageIdent::Lhs | StageIdent::Rhs => {
194            let m = tile_size.m() as usize;
195            let n = tile_size.n() as usize;
196            let k = tile_size.k() as usize;
197            let plane_dim = plane_dim as usize;
198            let k_local = k / plane_dim;
199
200            let shared_layout = comptime!(shared.layout);
201            let vector_size = ES::value();
202
203            let unit_id = UNIT_POS_X as usize;
204            let k_offset = k_local * unit_id;
205
206            let (strided_dim_count, contiguous_dim_count) = match (shared_layout, ident) {
207                (MatrixLayout::RowMajor, StageIdent::Lhs) => (m, k_local),
208                (MatrixLayout::RowMajor, StageIdent::Rhs) => (k_local, n),
209                (MatrixLayout::ColMajor, StageIdent::Lhs) => (k_local, m),
210                (MatrixLayout::ColMajor, StageIdent::Rhs) => (n, k_local),
211                _ => unreachable!(),
212            };
213
214            let (strided_dim_offset, contiguous_dim_offset) = match (shared_layout, ident) {
215                (MatrixLayout::RowMajor, StageIdent::Lhs)
216                | (MatrixLayout::ColMajor, StageIdent::Rhs) => (0, k_offset / vector_size),
217                (MatrixLayout::RowMajor, StageIdent::Rhs)
218                | (MatrixLayout::ColMajor, StageIdent::Lhs) => (k_offset, 0),
219                _ => unreachable!(),
220            };
221
222            assert!(contiguous_dim_count % vector_size == 0);
223            let vector_count_in_dim = contiguous_dim_count / vector_size;
224
225            for i in 0..strided_dim_count {
226                for j in 0..vector_count_in_dim {
227                    let vector = Vector::<N, ES>::cast_from(shared.get_vector(
228                        (i + strided_dim_offset) as u32,
229                        (j + contiguous_dim_offset) as u32,
230                    ));
231                    let vector_start = i * contiguous_dim_count + j * vector_size;
232                    for l in 0..vector_size {
233                        arr[vector_start + l] = vector.extract(l);
234                    }
235                }
236            }
237        }
238        StageIdent::Acc => {
239            panic!("Not yet implemented: Interleaved acc load from shared");
240        }
241        _ => panic!("Invalid ident for Interleaved load"),
242    }
243}
244
245#[cube]
246pub fn interleaved_load_zeros<N: Numeric>(arr: &mut Array<N>, #[comptime] tile_size: TileSize) {
247    let m = tile_size.m() as usize;
248    let n = tile_size.n() as usize;
249    let size = m * n;
250    for i in 0..size {
251        arr[i] = N::from_int(0);
252    }
253}
254
255#[cube]
256pub fn interleaved_write_to_shared<E: Numeric, ES: Size, A: Numeric>(
257    shared: &mut SharedTile<E>,
258    arr: &Array<A>,
259    #[comptime] tile_size: TileSize,
260) {
261    let mut shared = shared.view::<ES>();
262    let shared = &mut shared;
263    let m = tile_size.m();
264    let n = tile_size.n();
265    let out_vector_size = shared.container.vector_size().comptime() as u32;
266    let size_mn = m * n;
267
268    // `plane_sum` reduces across the plane, so every unit must participate. Only unit 0 stores.
269    #[unroll]
270    for i in 0..size_mn / out_vector_size {
271        let mut vector = Vector::<A, ES>::empty();
272        #[unroll]
273        for j in 0..out_vector_size {
274            vector.insert(
275                j as usize,
276                plane_sum(arr[(i * out_vector_size + j) as usize]),
277            );
278        }
279        if UNIT_POS_X == 0 {
280            let offs = shared.stage_offset(i);
281            shared.container[offs as usize] = Vector::cast_from(vector);
282        }
283    }
284}