Skip to main content

teeny_kernels/nn/fused/
conv2d_bn_silu.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#![allow(non_snake_case)]
18
19use core::ops::BitAnd;
20
21use teeny_macros::kernel;
22use teeny_triton::triton::{
23    types::{AddOffsets, Comparison, Tensor},
24    *,
25};
26
27/// Fused Conv2d + BatchNorm2d (inference) + SiLU forward pass.
28///
29/// Epilog fusion: after the conv accumulation loop, applies BN affine and
30/// SiLU in registers before the final global store, eliminating 2 intermediate
31/// global memory round-trips vs 3 separate kernels.
32///
33/// BN parameters must be precomputed by the caller as:
34///   `bn_scale[c] = gamma[c] / sqrt(var[c] + eps)`
35///   `bn_shift[c] = beta[c] - bn_scale[c] * mean[c]`
36///
37/// Grid: `pid = ((b * C_OUT + c_out) * OH + oh) * num_ow_tiles + ow_tile`
38///
39/// Inference-only; no backward pass.
40#[kernel]
41pub fn conv2d_bn_silu_forward<
42    T: Triton,
43    const KH: i32,
44    const KW: i32,
45    const STRIDE_H: i32,
46    const STRIDE_W: i32,
47    const PAD_H: i32,
48    const PAD_W: i32,
49    const G: i32,
50    const BLOCK_OW: i32,
51>(
52    x_ptr: T::Pointer<f32>,
53    w_ptr: T::Pointer<f32>,
54    bn_scale_ptr: T::Pointer<f32>,
55    bn_shift_ptr: T::Pointer<f32>,
56    y_ptr: T::Pointer<f32>,
57    _B: i32,
58    C_IN: i32,
59    C_OUT: i32,
60    H: i32,
61    W: i32,
62    OH: i32,
63    OW: i32,
64) where
65    T::I32Tensor: Tensor<i32, 1>,
66    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
67    T::BoolTensor: BitAnd<Output = T::BoolTensor>,
68    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
69{
70    let pid = T::program_id(Axis::X);
71    let num_ow_tiles = T::cdiv(OW, BLOCK_OW);
72
73    // Decode flat pid → (b, c_out, oh, ow_tile).
74    let ow_tile = pid % num_ow_tiles;
75    let bco = pid / num_ow_tiles;
76    let oh = bco % OH;
77    let bc = bco / OH;
78    let c_out = bc % C_OUT;
79    let b = bc / C_OUT;
80
81    let ow_start = ow_tile * BLOCK_OW;
82    let ow_range = T::arange(0, BLOCK_OW) + ow_start;
83    let ow_mask = ow_range.lt(OW);
84
85    let out_bc_base = (b * C_OUT + c_out) * OH * OW;
86
87    let c_in_per_group = C_IN / G;
88    let g_idx = c_out / (C_OUT / G);
89    let c_in_start = g_idx * c_in_per_group;
90
91    // ── Conv accumulation (same as conv2d_forward) ────────────────────────────
92    let mut acc = T::zeros::<f32>(&[BLOCK_OW]);
93
94    let loop_bound = c_in_per_group * KH * KW;
95    for idx in 0..loop_bound {
96        let kw = idx % KW;
97        let kh_cin = idx / KW;
98        let kh = kh_cin % KH;
99        let c_in_local = kh_cin / KH;
100        let c_in = c_in_start + c_in_local;
101
102        let ih = oh * STRIDE_H + kh - PAD_H;
103        let iw_range = ow_range * STRIDE_W + kw - PAD_W;
104
105        #[allow(clippy::erasing_op)]
106        let ih_t = ow_range * 0 + ih;
107        let h_in_bounds = ih_t.ge(0) & ih_t.lt(H);
108        let w_in_bounds = iw_range.ge(0) & iw_range.lt(W);
109        let load_mask = ow_mask & h_in_bounds & w_in_bounds;
110
111        let x_offsets = iw_range + ((b * C_IN + c_in) * H * W + ih * W);
112        let x_tile = T::load(
113            x_ptr.add_offsets(x_offsets),
114            Some(load_mask),
115            Some(T::zeros::<f32>(&[BLOCK_OW])),
116            &[],
117            None,
118            None,
119            None,
120            false,
121        );
122
123        // Weight layout [C_OUT, C_IN/G, KH, KW]: load scalar and broadcast.
124        let w_idx = ((c_out * c_in_per_group + c_in_local) * KH + kh) * KW + kw;
125        let w_off = T::arange(0, 1) + w_idx;
126        let w_1 = T::load(
127            w_ptr.add_offsets(w_off),
128            None,
129            None,
130            &[],
131            None,
132            None,
133            None,
134            false,
135        );
136        let w_tile = T::broadcast_to(w_1, &[BLOCK_OW]);
137
138        acc = acc + x_tile * w_tile;
139    }
140
141    // ── BatchNorm epilog: acc = bn_scale[c_out] * acc + bn_shift[c_out] ───────
142    let bn_off = T::arange(0, 1) + c_out;
143    let scale_1 = T::load(
144        bn_scale_ptr.add_offsets(bn_off),
145        None,
146        None,
147        &[],
148        None,
149        None,
150        None,
151        false,
152    );
153    let scale = T::broadcast_to(scale_1, &[BLOCK_OW]);
154    let shift_1 = T::load(
155        bn_shift_ptr.add_offsets(bn_off),
156        None,
157        None,
158        &[],
159        None,
160        None,
161        None,
162        false,
163    );
164    let shift = T::broadcast_to(shift_1, &[BLOCK_OW]);
165    let bn_out = scale * acc + shift;
166
167    // ── SiLU epilog: y = x * sigmoid(x) = x / (1 + exp(-x)) ─────────────────
168    let one = T::full(&[BLOCK_OW], 1.0_f32);
169    let neg1 = T::full(&[BLOCK_OW], -1.0_f32);
170    let y = bn_out * (one / (one + T::exp(neg1 * bn_out)));
171
172    let out_offsets = ow_range + (out_bc_base + oh * OW);
173    T::store(
174        y_ptr.add_offsets(out_offsets),
175        y,
176        Some(ow_mask),
177        &[],
178        None,
179        None,
180    );
181}
182
183// ── RuntimeOp ────────────────────────────────────────────────────────────────
184//
185// Params layout: [weight [C_OUT, C_IN/G, KH, KW], bn_scale [C_OUT], bn_shift [C_OUT]]
186// pack_args order: x_ptr, w_ptr, bn_scale_ptr, bn_shift_ptr, y_ptr,
187//                  B, C_IN, C_OUT, H, W, OH, OW
188
189impl teeny_core::model::RuntimeOp for Conv2dBnSiluForward {
190    fn n_activation_inputs(&self) -> usize {
191        1
192    }
193
194    fn param_shapes(&self, input_shapes: &[&[usize]], output_shape: &[usize]) -> Vec<Vec<usize>> {
195        let c_in = input_shapes[0][1];
196        let c_out = output_shape[1];
197        vec![
198            vec![
199                c_out,
200                c_in / self.g as usize,
201                self.kh as usize,
202                self.kw as usize,
203            ],
204            vec![c_out],
205            vec![c_out],
206        ]
207    }
208
209    fn param_names(&self) -> &'static [&'static str] {
210        &["weight", "bn_scale", "bn_shift"]
211    }
212
213    fn pack_args(
214        &self,
215        inputs: &[(teeny_core::model::RawPtr, &[usize])],
216        params: &[teeny_core::model::RawPtr],
217        output: teeny_core::model::RawPtr,
218        output_shape: &[usize],
219        _output_row_stride: i32,
220        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
221    ) {
222        let input_shape = inputs[0].1;
223        visitor.visit_ptr(inputs[0].0); // x_ptr
224        visitor.visit_ptr(params[0]); // w_ptr
225        visitor.visit_ptr(params[1]); // bn_scale_ptr
226        visitor.visit_ptr(params[2]); // bn_shift_ptr
227        visitor.visit_ptr(output); // y_ptr
228        visitor.visit_i32(input_shape[0] as i32); // B
229        visitor.visit_i32(input_shape[1] as i32); // C_IN
230        visitor.visit_i32(output_shape[1] as i32); // C_OUT
231        visitor.visit_i32(input_shape[2] as i32); // H
232        visitor.visit_i32(input_shape[3] as i32); // W
233        visitor.visit_i32(output_shape[2] as i32); // OH
234        visitor.visit_i32(output_shape[3] as i32); // OW
235    }
236
237    fn block(&self) -> [u32; 3] {
238        [128, 1, 1]
239    }
240
241    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
242        let num_ow_tiles = output_shape[3].div_ceil(self.block_ow as usize);
243        [
244            (output_shape[0] * output_shape[1] * output_shape[2] * num_ow_tiles) as u32,
245            1,
246            1,
247        ]
248    }
249}