1#![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#[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#[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}