Skip to main content

teeny_kernels/nn/pad/
constant_pad2d.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, BitOr};
20
21use teeny_core::dtype::Num;
22use teeny_macros::kernel;
23use teeny_triton::triton::{
24    types::{AddOffsets, Comparison, Tensor},
25    *,
26};
27
28/// 2-D constant padding forward pass.
29///
30/// Grid: `pid = ((b * C + c) * OH + oh) * num_ow_tiles + ow_tile`
31///
32/// `OH = PT + H + PB`, `OW = PL + W + PR`.
33/// Positions outside input region are filled with `value`.
34#[kernel]
35pub fn constant_pad2d_forward<
36    T: Triton,
37    D: Num,
38    const PT: i32,
39    const PB: i32,
40    const PL: i32,
41    const PR: i32,
42    const BLOCK_OW: i32,
43>(
44    input_ptr: T::Pointer<D>,
45    output_ptr: T::Pointer<D>,
46    _B: i32,
47    C: i32,
48    H: i32,
49    W: i32,
50    OH: i32,
51    OW: i32,
52    value: f32,
53) where
54    T::I32Tensor: Tensor<i32, 1>,
55    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
56    T::BoolTensor: BitAnd<Output = T::BoolTensor>,
57    T::BoolTensor: BitOr<Output = T::BoolTensor>,
58    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
59{
60    let pid = T::program_id(Axis::X);
61    let num_ow_tiles = T::cdiv(OW, BLOCK_OW);
62
63    let ow_tile = pid % num_ow_tiles;
64    let rest = pid / num_ow_tiles;
65    let oh = rest % OH;
66    let bc = rest / OH;
67    let c = bc % C;
68    let b = bc / C;
69
70    let ow_start = ow_tile * BLOCK_OW;
71    let ow_range = T::arange(0, BLOCK_OW) + ow_start;
72    let ow_mask = ow_range.lt(OW);
73
74    let ih = oh - PT;
75    let value_vec = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_OW], value), None, false);
76    let out_offsets = ow_range + ((b * C + c) * OH + oh) * OW;
77
78    if ih < 0 || ih >= H {
79        T::store(
80            output_ptr.add_offsets(out_offsets),
81            value_vec,
82            Some(ow_mask),
83            &[],
84            None,
85            None,
86        );
87        return;
88    }
89
90    let iw_range = ow_range - PL;
91    let w_in_bounds = iw_range.ge(0) & iw_range.lt(W);
92    let combined_mask = ow_mask & w_in_bounds;
93    let in_bc_base = (b * C + c) * H * W + ih * W;
94
95    let tile = T::load(
96        input_ptr.add_offsets(iw_range + in_bc_base),
97        Some(combined_mask),
98        Some(value_vec),
99        &[],
100        None,
101        None,
102        None,
103        false,
104    );
105    let result = T::where_(combined_mask, tile, value_vec);
106    T::store(
107        output_ptr.add_offsets(out_offsets),
108        result,
109        Some(ow_mask),
110        &[],
111        None,
112        None,
113    );
114}
115
116/// 2-D constant padding backward pass.
117#[kernel]
118pub fn constant_pad2d_backward<
119    T: Triton,
120    D: Num,
121    const PT: i32,
122    const PB: i32,
123    const PL: i32,
124    const PR: i32,
125    const BLOCK_OW: i32,
126>(
127    dy_ptr: T::Pointer<D>,
128    dx_ptr: T::Pointer<D>,
129    _B: i32,
130    C: i32,
131    H: i32,
132    W: i32,
133    OH: i32,
134    OW: i32,
135) where
136    T::I32Tensor: Tensor<i32, 1>,
137    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
138    T::BoolTensor: BitAnd<Output = T::BoolTensor>,
139    T::BoolTensor: BitOr<Output = T::BoolTensor>,
140    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
141{
142    let pid = T::program_id(Axis::X);
143    let num_ow_tiles = T::cdiv(OW, BLOCK_OW);
144
145    let ow_tile = pid % num_ow_tiles;
146    let rest = pid / num_ow_tiles;
147    let oh = rest % OH;
148    let bc = rest / OH;
149    let c = bc % C;
150    let b = bc / C;
151
152    let ow_start = ow_tile * BLOCK_OW;
153    let ow_range = T::arange(0, BLOCK_OW) + ow_start;
154    let ow_mask = ow_range.lt(OW);
155
156    let ih = oh - PT;
157    if ih < 0 || ih >= H {
158        return;
159    }
160    let iw_range = ow_range - PL;
161
162    let w_in_bounds = iw_range.ge(0) & iw_range.lt(W);
163
164    let dy_bc_base = ((b * C + c) * OH + oh) * OW;
165    let dx_bc_base = (b * C + c) * H * W + ih * W;
166
167    let store_mask = ow_mask & w_in_bounds;
168
169    let dy_offsets = ow_range + dy_bc_base;
170    let dy_tile = T::load(
171        dy_ptr.add_offsets(dy_offsets),
172        Some(ow_mask),
173        Some(T::zeros::<D>(&[BLOCK_OW])),
174        &[],
175        None,
176        None,
177        None,
178        false,
179    );
180
181    let dx_offsets = iw_range + dx_bc_base;
182    T::store(
183        dx_ptr.add_offsets(dx_offsets),
184        dy_tile,
185        Some(store_mask),
186        &[],
187        None,
188        None,
189    );
190}
191
192pub struct ConstantPad2dOp<'a, T: Num> {
193    pub forward: ConstantPad2dForward<T>,
194    pub backward: ConstantPad2dBackward<T>,
195    _marker: core::marker::PhantomData<&'a ()>,
196}