1use super::{BasicOps, DType, Int, Tensor, TensorPrimitive, backend::Backend};
2use super::ops::PadMode;
3use super::module::{
4 avg_pool1d, avg_pool2d,
5 max_pool1d, max_pool2d, max_pool1d_with_indices, max_pool2d_with_indices,
6};
7
8pub(super) fn volume_planes<B: Backend, K: BasicOps<B>>(input: Tensor<B, 5, K>) -> Tensor<B, 4, K> {
9 let [batch, channels, depth, height, width] = input.dims();
10 let planes = batch.checked_mul(depth).expect("pooling plane count overflow");
11 input.permute([0, 2, 1, 3, 4]).reshape([planes, channels, height, width])
12}
13
14pub(super) fn plane_depth_lines<B: Backend, K: BasicOps<B>>(
15 planes: Tensor<B, 4, K>,
16 batch: usize,
17 depth: usize,
18) -> Tensor<B, 3, K> {
19 let [_, channels, height, width] = planes.dims();
20 let lines = batch.checked_mul(channels)
21 .and_then(|count| count.checked_mul(height))
22 .and_then(|count| count.checked_mul(width))
23 .expect("pooling depth line count overflow");
24 planes.reshape([batch, depth, channels, height, width])
25 .permute([0, 2, 3, 4, 1]).reshape([lines, 1, depth])
26}
27
28pub(super) fn depth_lines_volume<B: Backend, K: BasicOps<B>>(
29 lines: Tensor<B, 3, K>,
30 batch: usize,
31 channels: usize,
32 height: usize,
33 width: usize,
34) -> Tensor<B, 5, K> {
35 let [_, _, depth] = lines.dims();
36 lines.reshape([batch, channels, height, width, depth]).permute([0, 1, 4, 2, 3])
37}
38
39pub fn max_pool3d<B: Backend>(
45 input: Tensor<B, 5>,
46 kernel_size: [usize; 3],
47 stride: [usize; 3],
48 padding: [usize; 3],
49 dilation: [usize; 3],
50 ceil_mode: bool,
51) -> Tensor<B, 5> {
52 let [batch, channels, depth, _, _] = input.dims();
53 let planes = max_pool2d(
54 volume_planes(input),
55 [kernel_size[1], kernel_size[2]],
56 [stride[1], stride[2]],
57 [padding[1], padding[2]],
58 [dilation[1], dilation[2]],
59 ceil_mode,
60 );
61 let [_, _, height, width] = planes.dims();
62 let lines = max_pool1d(
63 plane_depth_lines(planes, batch, depth),
64 kernel_size[0], stride[0], padding[0], dilation[0], ceil_mode,
65 );
66 depth_lines_volume(lines, batch, channels, height, width)
67}
68
69pub fn max_pool3d_with_indices<B: Backend>(
76 input: Tensor<B, 5>,
77 kernel_size: [usize; 3],
78 stride: [usize; 3],
79 padding: [usize; 3],
80 dilation: [usize; 3],
81 ceil_mode: bool,
82) -> (Tensor<B, 5>, Tensor<B, 5, Int>) {
83 let [batch, channels, depth, input_height, input_width] = input.dims();
84 assert!(depth > 0 && input_height > 0 && input_width > 0,
85 "volume pooling indices require non-empty spatial axes");
86 let area = input_height.checked_mul(input_width).expect("pooling plane area overflow");
87 let volume = depth.checked_mul(area).expect("pooling volume size overflow");
88 assert!(volume <= i64::MAX as usize, "volume pooling indices exceed I64");
89 let (planes, plane_indices) = max_pool2d_with_indices(
90 volume_planes(input),
91 [kernel_size[1], kernel_size[2]],
92 [stride[1], stride[2]],
93 [padding[1], padding[2]],
94 [dilation[1], dilation[2]],
95 ceil_mode,
96 );
97 let [_, _, height, width] = planes.dims();
98 let (lines, depth_indices) = max_pool1d_with_indices(
99 plane_depth_lines(planes, batch, depth),
100 kernel_size[0], stride[0], padding[0], dilation[0], ceil_mode,
101 );
102 let depth_indices = depth_indices.cast(DType::I64);
103 let safe_depth_indices = depth_indices.clone().clamp(0, depth as i64 - 1);
104 let spatial_indices = plane_depth_lines(plane_indices.cast(DType::I64), batch, depth)
105 .gather(2, safe_depth_indices.clone());
106 let invalid = depth_indices.clone().lower_elem(0)
107 .bool_or(depth_indices.clone().greater_equal_elem(depth as i64))
108 .bool_or(spatial_indices.clone().lower_elem(0))
109 .bool_or(spatial_indices.clone().greater_equal_elem(area as i64));
110 let indices = (safe_depth_indices.mul_scalar(area as i64)
111 + spatial_indices.clamp(0, area as i64 - 1)).mask_fill(invalid, -1);
112 (
113 depth_lines_volume(lines, batch, channels, height, width),
114 depth_lines_volume(indices, batch, channels, height, width),
115 )
116}
117
118pub fn avg_pool3d<B: Backend>(
124 input: Tensor<B, 5>,
125 kernel_size: [usize; 3],
126 stride: [usize; 3],
127 padding: [usize; 3],
128 count_include_pad: bool,
129 ceil_mode: bool,
130) -> Tensor<B, 5> {
131 Tensor::new(TensorPrimitive::Float(B::avg_pool3d(input.primitive.tensor(), kernel_size,
132 stride, padding, count_include_pad, ceil_mode)))
133}
134
135pub fn adaptive_avg_pool3d<B: Backend>(
141 input: Tensor<B, 5>,
142 output_size: [usize; 3],
143) -> Tensor<B, 5> {
144 Tensor::new(TensorPrimitive::Float(B::adaptive_avg_pool3d(input.primitive.tensor(), output_size)))
145}
146
147fn average_excluding_explicit_padding<B: Backend, const D: usize, const N: usize>(
148 input: Tensor<B, D>,
149 padding: [(usize, usize); N],
150 pool: impl Fn(Tensor<B, D>) -> Tensor<B, D>,
151) -> Tensor<B, D> {
152 let storage = input.dtype();
153 let compute = if storage == DType::F64 { DType::F64 } else { DType::F32 };
154 let mut visible_shape = input.dims();
155 visible_shape[0] = 1;
156 visible_shape[1] = 1;
157 let visible = Tensor::<B, D>::ones(visible_shape, (&input.device(), compute))
158 .pad(padding, PadMode::Constant(0.0));
159 let values = pool(input.cast(compute).pad(padding, PadMode::Constant(0.0)));
160 let coverage = pool(visible);
161 (values / coverage).cast(storage)
162}
163
164pub fn avg_pool1d_padded<B: Backend>(
170 input: Tensor<B, 3>,
171 kernel_size: usize,
172 stride: usize,
173 padding: [(usize, usize); 1],
174 count_include_pad: bool,
175 ceil_mode: bool,
176) -> Tensor<B, 3> {
177 let [(left, right)] = padding;
178 if left == right {
179 return avg_pool1d(input, kernel_size, stride, left, count_include_pad, ceil_mode);
180 }
181 if count_include_pad {
182 return avg_pool1d(input.pad(padding, PadMode::Constant(0.0)),
183 kernel_size, stride, 0, true, ceil_mode);
184 }
185 average_excluding_explicit_padding(input, padding,
186 |input| avg_pool1d(input, kernel_size, stride, 0, true, ceil_mode))
187}
188
189pub fn avg_pool2d_padded<B: Backend>(
194 input: Tensor<B, 4>,
195 kernel_size: [usize; 2],
196 stride: [usize; 2],
197 padding: [(usize, usize); 2],
198 count_include_pad: bool,
199 ceil_mode: bool,
200) -> Tensor<B, 4> {
201 if padding.iter().all(|(start, end)| start == end) {
202 return avg_pool2d(input, kernel_size, stride, padding.map(|(start, _)| start),
203 count_include_pad, ceil_mode);
204 }
205 if count_include_pad {
206 return avg_pool2d(input.pad(padding, PadMode::Constant(0.0)),
207 kernel_size, stride, [0; 2], true, ceil_mode);
208 }
209 average_excluding_explicit_padding(input, padding,
210 |input| avg_pool2d(input, kernel_size, stride, [0; 2], true, ceil_mode))
211}
212
213pub fn avg_pool3d_padded<B: Backend>(
217 input: Tensor<B, 5>,
218 kernel_size: [usize; 3],
219 stride: [usize; 3],
220 padding: [(usize, usize); 3],
221 count_include_pad: bool,
222 ceil_mode: bool,
223) -> Tensor<B, 5> {
224 if padding.iter().all(|(start, end)| start == end) {
225 return avg_pool3d(input, kernel_size, stride, padding.map(|(start, _)| start),
226 count_include_pad, ceil_mode);
227 }
228 if count_include_pad {
229 return avg_pool3d(input.pad(padding, PadMode::Constant(0.0)),
230 kernel_size, stride, [0; 3], true, ceil_mode);
231 }
232 average_excluding_explicit_padding(input, padding,
233 |input| avg_pool3d(input, kernel_size, stride, [0; 3], true, ceil_mode))
234}
235
236pub fn max_pool1d_padded<B: Backend>(
238 input: Tensor<B, 3>,
239 kernel_size: usize,
240 stride: usize,
241 padding: [(usize, usize); 1],
242 dilation: usize,
243 ceil_mode: bool,
244) -> Tensor<B, 3> {
245 let [(left, right)] = padding;
246 let (input, padding) = if left == right {
247 (input, left)
248 } else {
249 (input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)), 0)
250 };
251 max_pool1d(input, kernel_size, stride, padding, dilation, ceil_mode)
252}
253
254pub fn max_pool2d_padded<B: Backend>(
256 input: Tensor<B, 4>,
257 kernel_size: [usize; 2],
258 stride: [usize; 2],
259 padding: [(usize, usize); 2],
260 dilation: [usize; 2],
261 ceil_mode: bool,
262) -> Tensor<B, 4> {
263 let (input, padding) = if padding.iter().all(|(start, end)| start == end) {
264 (input, padding.map(|(start, _)| start))
265 } else {
266 (input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)), [0; 2])
267 };
268 max_pool2d(input, kernel_size, stride, padding, dilation, ceil_mode)
269}
270
271pub fn max_pool3d_padded<B: Backend>(
273 input: Tensor<B, 5>,
274 kernel_size: [usize; 3],
275 stride: [usize; 3],
276 padding: [(usize, usize); 3],
277 dilation: [usize; 3],
278 ceil_mode: bool,
279) -> Tensor<B, 5> {
280 let (input, padding) = if padding.iter().all(|(start, end)| start == end) {
281 (input, padding.map(|(start, _)| start))
282 } else {
283 (input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)), [0; 3])
284 };
285 max_pool3d(input, kernel_size, stride, padding, dilation, ceil_mode)
286}
287
288fn unpad_pool_indices<B: Backend, const D: usize, const N: usize>(
289 indices: Tensor<B, D, Int>,
290 input_size: [usize; N],
291 padding: [(usize, usize); N],
292) -> Tensor<B, D, Int> {
293 let padded_size: [usize; N] = core::array::from_fn(|axis| {
294 input_size[axis].checked_add(padding[axis].0)
295 .and_then(|size| size.checked_add(padding[axis].1))
296 .expect("padded pooling index extent overflow")
297 });
298 let padded_volume = padded_size.iter().try_fold(1usize,
299 |size, axis| size.checked_mul(*axis)).expect("padded pooling index volume overflow");
300 let input_volume = input_size.iter().try_fold(1usize,
301 |size, axis| size.checked_mul(*axis)).expect("pooling input index volume overflow");
302 assert!(padded_volume > 0 && padded_volume <= i64::MAX as usize
303 && input_volume <= i64::MAX as usize, "pooling indices cannot be represented in I64");
304 let indices = indices.cast(DType::I64);
305 if input_volume == 0 {
306 return indices.zeros_like().sub_scalar(1);
307 }
308 let mut invalid = indices.clone().lower_elem(0)
309 .bool_or(indices.clone().greater_equal_elem(padded_volume as i64));
310 let mut remaining = indices.clone().clamp(0, padded_volume as i64 - 1);
311 let mut unpadded = indices.zeros_like();
312 let mut input_stride = 1usize;
313 for axis in (0..N).rev() {
314 let coordinate = remaining.clone().remainder_scalar(padded_size[axis] as i64)
315 .sub_scalar(padding[axis].0 as i64);
316 remaining = remaining.div_scalar(padded_size[axis] as i64);
317 invalid = invalid.bool_or(coordinate.clone().lower_elem(0))
318 .bool_or(coordinate.clone().greater_equal_elem(input_size[axis] as i64));
319 unpadded = unpadded + coordinate.clamp(0, input_size[axis].saturating_sub(1) as i64)
320 .mul_scalar(input_stride as i64);
321 input_stride = input_stride.checked_mul(input_size[axis])
322 .expect("pooling input index stride overflow");
323 }
324 unpadded.mask_fill(invalid, -1)
325}
326
327pub fn max_pool1d_with_indices_padded<B: Backend>(
333 input: Tensor<B, 3>,
334 kernel_size: usize,
335 stride: usize,
336 padding: [(usize, usize); 1],
337 dilation: usize,
338 ceil_mode: bool,
339) -> (Tensor<B, 3>, Tensor<B, 3, Int>) {
340 let [(left, right)] = padding;
341 if left == right {
342 return max_pool1d_with_indices(input, kernel_size, stride, left, dilation, ceil_mode);
343 }
344 let [_, _, length] = input.dims();
345 let (values, indices) = max_pool1d_with_indices(
346 input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)),
347 kernel_size, stride, 0, dilation, ceil_mode,
348 );
349 (values, unpad_pool_indices(indices, [length], padding))
350}
351
352pub fn max_pool2d_with_indices_padded<B: Backend>(
356 input: Tensor<B, 4>,
357 kernel_size: [usize; 2],
358 stride: [usize; 2],
359 padding: [(usize, usize); 2],
360 dilation: [usize; 2],
361 ceil_mode: bool,
362) -> (Tensor<B, 4>, Tensor<B, 4, Int>) {
363 if padding.iter().all(|(start, end)| start == end) {
364 return max_pool2d_with_indices(input, kernel_size, stride,
365 padding.map(|(start, _)| start), dilation, ceil_mode);
366 }
367 let [_, _, height, width] = input.dims();
368 let (values, indices) = max_pool2d_with_indices(
369 input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)),
370 kernel_size, stride, [0; 2], dilation, ceil_mode,
371 );
372 (values, unpad_pool_indices(indices, [height, width], padding))
373}
374
375pub fn max_pool3d_with_indices_padded<B: Backend>(
380 input: Tensor<B, 5>,
381 kernel_size: [usize; 3],
382 stride: [usize; 3],
383 padding: [(usize, usize); 3],
384 dilation: [usize; 3],
385 ceil_mode: bool,
386) -> (Tensor<B, 5>, Tensor<B, 5, Int>) {
387 if padding.iter().all(|(start, end)| start == end) {
388 return max_pool3d_with_indices(input, kernel_size, stride,
389 padding.map(|(start, _)| start), dilation, ceil_mode);
390 }
391 let [_, _, depth, height, width] = input.dims();
392 let (values, indices) = max_pool3d_with_indices(
393 input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)),
394 kernel_size, stride, [0; 3], dilation, ceil_mode,
395 );
396 (values, unpad_pool_indices(indices, [depth, height, width], padding))
397}