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 avgpool3d_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::zeros::<D>(&[BLOCK_OW]);
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::zeros::<D>(&[BLOCK_OW])),
96 &[],
97 None,
98 None,
99 None,
100 false,
101 );
102 acc = acc + tile;
103 }
104
105 let ksize_1 = T::full::<i32>(&[1], KD * KH * KW);
106 let ksize_f_1 = T::cast::<i32, D>(ksize_1, None, false);
107 let ksize = T::broadcast_to(ksize_f_1, &[BLOCK_OW]);
108 let result = acc / ksize;
109
110 let out_offsets = ow_range + out_base;
111 T::store(
112 output_ptr.add_offsets(out_offsets),
113 result,
114 Some(ow_mask),
115 &[],
116 None,
117 None,
118 );
119}
120
121#[kernel]
127pub fn avgpool3d_backward<
128 T: Triton,
129 D: Num,
130 const KD: i32,
131 const KH: i32,
132 const KW: i32,
133 const STRIDE_D: i32,
134 const STRIDE_H: i32,
135 const STRIDE_W: i32,
136 const BLOCK_OW: i32,
137>(
138 dy_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 dy_base = ((b * C + c) * OD * OH * OW) + od * OH * OW + oh * OW;
170 let dx_bc_base = (b * C + c) * Dv * H * W;
171
172 let dy_offsets = ow_range + dy_base;
173 let dy_tile = T::load(
174 dy_ptr.add_offsets(dy_offsets),
175 Some(ow_mask),
176 Some(T::zeros::<D>(&[BLOCK_OW])),
177 &[],
178 None,
179 None,
180 None,
181 false,
182 );
183 let ksize_1 = T::full::<i32>(&[1], KD * KH * KW);
184 let ksize_f_1 = T::cast::<i32, D>(ksize_1, None, false);
185 let ksize = T::broadcast_to(ksize_f_1, &[BLOCK_OW]);
186 let grad = dy_tile / ksize;
187
188 let loop_bound = KD * KH * KW;
189 for idx in 0..loop_bound {
190 let kw = idx % KW;
191 let tmp = idx / KW;
192 let kh = tmp % KH;
193 let kd = tmp / KH;
194
195 let id = od * STRIDE_D + kd;
196 let ih = oh * STRIDE_H + kh;
197 let iw_range = ow_range * STRIDE_W + kw;
198 let dx_offsets = iw_range + (dx_bc_base + id * H * W + ih * W);
199 T::atomic_add(
200 dx_ptr.add_offsets(dx_offsets),
201 grad,
202 Some(ow_mask),
203 None,
204 None,
205 );
206 }
207}
208
209pub struct Avgpool3dOp<'a, T: Num> {
210 pub forward: Avgpool3dForward<T>,
211 pub backward: Avgpool3dBackward<T>,
212 _marker: core::marker::PhantomData<&'a ()>,
213}