Skip to main content

teeny_kernels/nn/pool/
maxpool3d.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 teeny_core::dtype::Num;
20use teeny_macros::kernel;
21use teeny_triton::triton::{
22    types::{AddOffsets, Comparison, Tensor},
23    *,
24};
25
26/// 3-D max-pooling forward pass.
27///
28/// Grid: `pid = (((b * C + c) * OD + od) * OH + oh) * num_ow_tiles + ow_tile`
29///
30/// **Constraints**: no padding;
31/// `OD = (D - KD) / STRIDE_D + 1`, `OH = (H - KH) / STRIDE_H + 1`,
32/// `OW = (W - KW) / STRIDE_W + 1`.
33#[kernel]
34pub fn maxpool3d_forward<
35    T: Triton,
36    D: Num,
37    const KD: i32,
38    const KH: i32,
39    const KW: i32,
40    const STRIDE_D: i32,
41    const STRIDE_H: i32,
42    const STRIDE_W: i32,
43    const BLOCK_OW: i32,
44>(
45    input_ptr: T::Pointer<D>,
46    output_ptr: T::Pointer<D>,
47    _B: i32,
48    C: i32,
49    Dv: i32,
50    H: i32,
51    W: i32,
52    OD: i32,
53    OH: i32,
54    OW: i32,
55) where
56    T::I32Tensor: Tensor<i32, 1>,
57    T::I32Tensor: Comparison<i32, BoolTensor = 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 rest2 = rest / OH;
67    let od = rest2 % OD;
68    let bco = rest2 / OD;
69    let c = bco % C;
70    let b = bco / C;
71
72    let ow_start = ow_tile * BLOCK_OW;
73    let ow_range = T::arange(0, BLOCK_OW) + ow_start;
74    let ow_mask = ow_range.lt(OW);
75
76    let in_bc_base = (b * C + c) * Dv * H * W;
77    let out_base = ((b * C + c) * OD * OH * OW) + od * OH * OW + oh * OW;
78
79    let mut acc = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_OW], -3.4028235e38_f32), None, false);
80
81    let loop_bound = KD * KH * KW;
82    for idx in 0..loop_bound {
83        let kw = idx % KW;
84        let tmp = idx / KW;
85        let kh = tmp % KH;
86        let kd = tmp / KH;
87
88        let id = od * STRIDE_D + kd;
89        let ih = oh * STRIDE_H + kh;
90        let iw_range = ow_range * STRIDE_W + kw;
91        let in_offsets = iw_range + (in_bc_base + id * H * W + ih * W);
92        let tile = T::load(
93            input_ptr.add_offsets(in_offsets),
94            Some(ow_mask),
95            Some(T::cast::<f32, D>(
96                T::full::<f32>(&[BLOCK_OW], -3.4028235e38_f32),
97                None,
98                false,
99            )),
100            &[],
101            None,
102            None,
103            None,
104            false,
105        );
106        acc = T::maximum(acc, tile);
107    }
108
109    let out_offsets = ow_range + out_base;
110    T::store(
111        output_ptr.add_offsets(out_offsets),
112        acc,
113        Some(ow_mask),
114        &[],
115        None,
116        None,
117    );
118}
119
120/// 3-D max-pooling backward pass.
121///
122/// Re-scans the input window and scatters `dy` to positions where
123/// `input == output_max`. `dx` must be zero-initialised before launch.
124#[kernel]
125pub fn maxpool3d_backward<
126    T: Triton,
127    D: Num,
128    const KD: i32,
129    const KH: i32,
130    const KW: i32,
131    const STRIDE_D: i32,
132    const STRIDE_H: i32,
133    const STRIDE_W: i32,
134    const BLOCK_OW: i32,
135>(
136    dy_ptr: T::Pointer<D>,
137    x_ptr: T::Pointer<D>,
138    y_ptr: T::Pointer<D>,
139    dx_ptr: T::Pointer<D>,
140    _B: i32,
141    C: i32,
142    Dv: i32,
143    H: i32,
144    W: i32,
145    OD: i32,
146    OH: i32,
147    OW: i32,
148) where
149    T::I32Tensor: Tensor<i32, 1>,
150    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
151    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
152{
153    let pid = T::program_id(Axis::X);
154    let num_ow_tiles = T::cdiv(OW, BLOCK_OW);
155
156    let ow_tile = pid % num_ow_tiles;
157    let rest = pid / num_ow_tiles;
158    let oh = rest % OH;
159    let rest2 = rest / OH;
160    let od = rest2 % OD;
161    let bco = rest2 / OD;
162    let c = bco % C;
163    let b = bco / C;
164
165    let ow_start = ow_tile * BLOCK_OW;
166    let ow_range = T::arange(0, BLOCK_OW) + ow_start;
167    let ow_mask = ow_range.lt(OW);
168
169    let in_bc_base = (b * C + c) * Dv * H * W;
170    let out_base = ((b * C + c) * OD * OH * OW) + od * OH * OW + oh * OW;
171
172    let out_offsets = ow_range + out_base;
173    let dy_tile = T::load(
174        dy_ptr.add_offsets(out_offsets),
175        Some(ow_mask),
176        Some(T::zeros::<D>(&[BLOCK_OW])),
177        &[],
178        None,
179        None,
180        None,
181        false,
182    );
183    let y_tile = T::load(
184        y_ptr.add_offsets(out_offsets),
185        Some(ow_mask),
186        Some(T::zeros::<D>(&[BLOCK_OW])),
187        &[],
188        None,
189        None,
190        None,
191        false,
192    );
193
194    let loop_bound = KD * KH * KW;
195    for idx in 0..loop_bound {
196        let kw = idx % KW;
197        let tmp = idx / KW;
198        let kh = tmp % KH;
199        let kd = tmp / KH;
200
201        let id = od * STRIDE_D + kd;
202        let ih = oh * STRIDE_H + kh;
203        let iw_range = ow_range * STRIDE_W + kw;
204        let in_offsets = iw_range + (in_bc_base + id * H * W + ih * W);
205        let x_tile = T::load(
206            x_ptr.add_offsets(in_offsets),
207            Some(ow_mask),
208            Some(T::cast::<f32, D>(
209                T::full::<f32>(&[BLOCK_OW], -3.4028235e38_f32),
210                None,
211                false,
212            )),
213            &[],
214            None,
215            None,
216            None,
217            false,
218        );
219        let is_max = T::eq(x_tile, y_tile);
220        let grad = T::where_(is_max, dy_tile, T::zeros::<D>(&[BLOCK_OW]));
221        T::atomic_add(
222            dx_ptr.add_offsets(in_offsets),
223            grad,
224            Some(ow_mask),
225            None,
226            None,
227        );
228    }
229}
230
231pub struct Maxpool3dOp<'a, T: Num> {
232    pub forward: Maxpool3dForward<T>,
233    pub backward: Maxpool3dBackward<T>,
234    _marker: core::marker::PhantomData<&'a ()>,
235}