Skip to main content

teeny_kernels/math/
gemm.rs

1/*
2 * Copyright (c) 2026 Teenygrad.
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 *   http://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17//! 2-D matrix multiply (MatMul / Gemm) Triton kernels.
18//!
19//! Tiled GEMM using `T::make_tensor_descriptor` + `T::dot` for Tensor Core
20//! utilisation — one CTA computes one `[BLOCK_M, BLOCK_N]` (or `[BLOCK_M,
21//! BLOCK_K]` / `[BLOCK_K, BLOCK_N]` for the backward kernels) output tile,
22//! accumulating over `K`/`N`/`M`-tiles with `T::dot` rather than a scalar
23//! multiply-and-reduce per element. Same swizzled-pid / tensor-descriptor
24//! structure as [`crate::nn::mlp::linear`]'s `linear_forward`/`linear_backward`.
25//!
26//! Grid: one CTA per output tile. Block: `[128, 1, 1]`.
27
28#![allow(non_snake_case)]
29
30use teeny_core::dtype::Num;
31use teeny_macros::kernel;
32use teeny_triton::triton::{PaddingOption, *};
33
34// ── MatMul Forward ────────────────────────────────────────────────────────────
35//
36// C[M, N] = A[M, K] @ B[K, N]
37//
38// Kernel grid: one CTA per [BLOCK_M, BLOCK_N] output tile, pids swizzled by
39// GROUP_M for L2 locality (same scheme as linear_forward).
40
41/// Forward: C = A @ B
42// ANCHOR: matmul_forward
43#[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// ANCHOR_END: matmul_forward
108
109/// Backward: dA = dC @ B^T
110///
111/// Grid: one CTA per `[BLOCK_M, BLOCK_K]` tile of dA.
112/// dA\[m, k\] = sum_n dC\[m, n\] * B\[k, n\]  (B\[k, n\] = B^T\[n, k\])
113#[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/// Backward: dB = A^T @ dC
169///
170/// Grid: one CTA per `[BLOCK_K, BLOCK_N]` tile of dB.
171/// dB\[k, n\] = sum_m A\[m, k\] * dC\[m, n\]
172#[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
227// ── RuntimeOp for MatMul ──────────────────────────────────────────────────────
228
229/// Swizzle group size for `matmul_forward`'s pid decomposition — see `linear_forward`.
230const 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        // A: [M, K], B: [K, N], C: [M, N]
274        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); // a_ptr
278        visitor.visit_ptr(inputs[1].0); // b_ptr
279        visitor.visit_ptr(output); // c_ptr
280        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    // For backward we pack args for dA kernel. The lowering handles dB separately.
303    // This is a simplified backward that only handles dA = dC @ B^T.
304    #[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); // dc_ptr
321        visitor.visit_ptr(inputs[1].0); // b_ptr
322        visitor.visit_ptr(grad_inputs[0]); // da_ptr
323        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}