cubek_std/tile/variants/instruction/
interleaved.rs1use cubecl::prelude::*;
2
3use crate::{
4 MatrixLayout, StageIdent, TileSize,
5 tile::{SharedTile, Tile, TileKind, TileKindExpand, TileScope},
6};
7
8#[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 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 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#[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 #[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}