1#![allow(non_snake_case)]
29
30use teeny_core::dtype::Num;
31use teeny_macros::kernel;
32use teeny_triton::triton::{PaddingOption, *};
33
34#[kernel]
44pub fn matmul_forward<
45 T: Triton,
46 D: Num,
47 const BLOCK_M: i32,
48 const BLOCK_N: i32,
49 const BLOCK_K: i32,
50 const GROUP_M: i32,
51>(
52 a_ptr: T::Pointer<D>,
53 b_ptr: T::Pointer<D>,
54 c_ptr: T::Pointer<D>,
55 M: i32,
56 N: i32,
57 K: i32,
58) {
59 let pid = T::program_id(Axis::X);
60 let num_pid_m = T::cdiv(M, BLOCK_M);
61 let num_pid_n = T::cdiv(N, BLOCK_N);
62 let num_pid_in_group = GROUP_M * num_pid_n;
63 let group_id = pid / num_pid_in_group;
64 let first_pid_m = group_id * GROUP_M;
65 let remaining_m = num_pid_m - first_pid_m;
66 let group_size_m = if remaining_m < GROUP_M {
67 remaining_m
68 } else {
69 GROUP_M
70 };
71 let pid_in_group = pid % num_pid_in_group;
72 let pid_m = first_pid_m + (pid_in_group % group_size_m);
73 let pid_n = pid_in_group / group_size_m;
74
75 let a_desc = T::make_tensor_descriptor(
76 a_ptr,
77 &[M, K],
78 &[K, 1],
79 &[BLOCK_M, BLOCK_K],
80 Some(PaddingOption::Zero),
81 );
82 let b_desc = T::make_tensor_descriptor(
83 b_ptr,
84 &[K, N],
85 &[N, 1],
86 &[BLOCK_K, BLOCK_N],
87 Some(PaddingOption::Zero),
88 );
89
90 let mut acc = T::zeros::<D>(&[BLOCK_M, BLOCK_N]);
91 let k_tiles = T::cdiv(K, BLOCK_K);
92 for k in 0..k_tiles {
93 let a = T::load_tensor_descriptor(a_desc, &[pid_m * BLOCK_M, k * BLOCK_K]);
94 let b = T::load_tensor_descriptor(b_desc, &[k * BLOCK_K, pid_n * BLOCK_N]);
95 acc = T::dot::<D, D>(a, b, Some(acc), InputPrecision::TF32, None);
96 }
97
98 let c_desc = T::make_tensor_descriptor(
99 c_ptr,
100 &[M, N],
101 &[N, 1],
102 &[BLOCK_M, BLOCK_N],
103 Some(PaddingOption::Zero),
104 );
105 T::store_tensor_descriptor(c_desc, &[pid_m * BLOCK_M, pid_n * BLOCK_N], acc);
106}
107#[kernel]
114pub fn matmul_backward_da<
115 T: Triton,
116 D: Num,
117 const BLOCK_M: i32,
118 const BLOCK_N: i32,
119 const BLOCK_K: i32,
120 const GROUP_M: i32,
121>(
122 dc_ptr: T::Pointer<D>,
123 b_ptr: T::Pointer<D>,
124 da_ptr: T::Pointer<D>,
125 M: i32,
126 N: i32,
127 K: i32,
128) {
129 let pid = T::program_id(Axis::X);
130 let num_pid_k = T::cdiv(K, BLOCK_K);
131 let pid_k = pid % num_pid_k;
132 let pid_m = pid / num_pid_k;
133
134 let dc_desc = T::make_tensor_descriptor(
135 dc_ptr,
136 &[M, N],
137 &[N, 1],
138 &[BLOCK_M, BLOCK_N],
139 Some(PaddingOption::Zero),
140 );
141 let b_desc = T::make_tensor_descriptor(
142 b_ptr,
143 &[K, N],
144 &[N, 1],
145 &[BLOCK_K, BLOCK_N],
146 Some(PaddingOption::Zero),
147 );
148
149 let mut acc = T::zeros::<D>(&[BLOCK_M, BLOCK_K]);
150 let n_tiles = T::cdiv(N, BLOCK_N);
151 for n in 0..n_tiles {
152 let dc = T::load_tensor_descriptor(dc_desc, &[pid_m * BLOCK_M, n * BLOCK_N]);
153 let b = T::load_tensor_descriptor(b_desc, &[pid_k * BLOCK_K, n * BLOCK_N]);
154 let b_t = T::trans(b, &[1, 0]);
155 acc = T::dot::<D, D>(dc, b_t, Some(acc), InputPrecision::TF32, None);
156 }
157
158 let da_desc = T::make_tensor_descriptor(
159 da_ptr,
160 &[M, K],
161 &[K, 1],
162 &[BLOCK_M, BLOCK_K],
163 Some(PaddingOption::Zero),
164 );
165 T::store_tensor_descriptor(da_desc, &[pid_m * BLOCK_M, pid_k * BLOCK_K], acc);
166}
167
168#[kernel]
173pub fn matmul_backward_db<
174 T: Triton,
175 D: Num,
176 const BLOCK_M: i32,
177 const BLOCK_N: i32,
178 const BLOCK_K: i32,
179 const GROUP_M: i32,
180>(
181 dc_ptr: T::Pointer<D>,
182 a_ptr: T::Pointer<D>,
183 db_ptr: T::Pointer<D>,
184 M: i32,
185 N: i32,
186 K: i32,
187) {
188 let pid = T::program_id(Axis::X);
189 let num_pid_n = T::cdiv(N, BLOCK_N);
190 let pid_n = pid % num_pid_n;
191 let pid_k = pid / num_pid_n;
192
193 let dc_desc = T::make_tensor_descriptor(
194 dc_ptr,
195 &[M, N],
196 &[N, 1],
197 &[BLOCK_M, BLOCK_N],
198 Some(PaddingOption::Zero),
199 );
200 let a_desc = T::make_tensor_descriptor(
201 a_ptr,
202 &[M, K],
203 &[K, 1],
204 &[BLOCK_M, BLOCK_K],
205 Some(PaddingOption::Zero),
206 );
207
208 let mut acc = T::zeros::<D>(&[BLOCK_K, BLOCK_N]);
209 let m_tiles = T::cdiv(M, BLOCK_M);
210 for m in 0..m_tiles {
211 let dc = T::load_tensor_descriptor(dc_desc, &[m * BLOCK_M, pid_n * BLOCK_N]);
212 let a = T::load_tensor_descriptor(a_desc, &[m * BLOCK_M, pid_k * BLOCK_K]);
213 let a_t = T::trans(a, &[1, 0]);
214 acc = T::dot::<D, D>(a_t, dc, Some(acc), InputPrecision::TF32, None);
215 }
216
217 let db_desc = T::make_tensor_descriptor(
218 db_ptr,
219 &[K, N],
220 &[N, 1],
221 &[BLOCK_K, BLOCK_N],
222 Some(PaddingOption::Zero),
223 );
224 T::store_tensor_descriptor(db_desc, &[pid_k * BLOCK_K, pid_n * BLOCK_N], acc);
225}
226
227const GROUP_M: i32 = 8;
231
232pub struct MatMulRuntimeOp<D: Num + Send + Sync + 'static> {
233 pub fwd_kernel: MatmulForward<D>,
234 pub bwd_da_kernel: MatmulBackwardDa<D>,
235 pub bwd_db_kernel: MatmulBackwardDb<D>,
236}
237
238impl<D: Num + Send + Sync + 'static> MatMulRuntimeOp<D> {
239 pub fn new(block_m: i32, block_n: i32, block_k: i32) -> Self {
240 Self {
241 fwd_kernel: MatmulForward::<D>::new(block_m, block_n, block_k, GROUP_M),
242 bwd_da_kernel: MatmulBackwardDa::<D>::new(block_m, block_n, block_k, GROUP_M),
243 bwd_db_kernel: MatmulBackwardDb::<D>::new(block_m, block_n, block_k, GROUP_M),
244 }
245 }
246
247 pub fn forward_source(&self) -> &str {
248 &self.fwd_kernel.source
249 }
250 pub fn kernel_name(&self) -> &str {
251 self.fwd_kernel.name
252 }
253}
254
255impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for MatMulRuntimeOp<D> {
256 fn n_activation_inputs(&self) -> usize {
257 2
258 }
259
260 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
261 vec![]
262 }
263
264 fn pack_args(
265 &self,
266 inputs: &[(teeny_core::model::RawPtr, &[usize])],
267 _: &[teeny_core::model::RawPtr],
268 output: teeny_core::model::RawPtr,
269 output_shape: &[usize],
270 _: i32,
271 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
272 ) {
273 let m = inputs[0].1.first().copied().unwrap_or(1) as i32;
275 let k = inputs[0].1.last().copied().unwrap_or(1) as i32;
276 let n = output_shape.last().copied().unwrap_or(1) as i32;
277 visitor.visit_ptr(inputs[0].0); visitor.visit_ptr(inputs[1].0); visitor.visit_ptr(output); visitor.visit_i32(m);
281 visitor.visit_i32(n);
282 visitor.visit_i32(k);
283 }
284
285 fn block(&self) -> [u32; 3] {
286 [128, 1, 1]
287 }
288
289 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
290 let m = output_shape.first().copied().unwrap_or(1) as u32;
291 let n = output_shape.last().copied().unwrap_or(1) as u32;
292 let pm = m.div_ceil(self.fwd_kernel.block_m as u32);
293 let pn = n.div_ceil(self.fwd_kernel.block_n as u32);
294 [pm * pn, 1, 1]
295 }
296
297 #[cfg(feature = "training")]
298 fn has_backward(&self) -> bool {
299 true
300 }
301
302 #[cfg(feature = "training")]
305 fn pack_backward_args(
306 &self,
307 inputs: &[(teeny_core::model::RawPtr, &[usize])],
308 _: &[teeny_core::model::RawPtr],
309 _: teeny_core::model::RawPtr,
310 output_shape: &[usize],
311 grad_output: teeny_core::model::RawPtr,
312 _: i32,
313 grad_inputs: &[teeny_core::model::RawPtr],
314 _: &[teeny_core::model::RawPtr],
315 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
316 ) {
317 let m = inputs[0].1.first().copied().unwrap_or(1) as i32;
318 let k = inputs[0].1.last().copied().unwrap_or(1) as i32;
319 let n = output_shape.last().copied().unwrap_or(1) as i32;
320 visitor.visit_ptr(grad_output); visitor.visit_ptr(inputs[1].0); visitor.visit_ptr(grad_inputs[0]); visitor.visit_i32(m);
324 visitor.visit_i32(n);
325 visitor.visit_i32(k);
326 }
327
328 #[cfg(feature = "training")]
329 fn backward_block(&self) -> [u32; 3] {
330 [128, 1, 1]
331 }
332
333 #[cfg(feature = "training")]
334 fn backward_grid(&self, input_shapes: &[&[usize]], _: &[usize]) -> [u32; 3] {
335 let m = input_shapes
336 .first()
337 .and_then(|s| s.first())
338 .copied()
339 .unwrap_or(1) as u32;
340 let k = input_shapes
341 .first()
342 .and_then(|s| s.last())
343 .copied()
344 .unwrap_or(1) as u32;
345 let pm = m.div_ceil(self.bwd_da_kernel.block_m as u32);
346 let pk = k.div_ceil(self.bwd_da_kernel.block_k as u32);
347 [pm * pk, 1, 1]
348 }
349}