Skip to main content

ruda_tensor/api/
spatial_pool.rs

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
39/// Pool native `[batch, channels, depth, height, width]` activations by maximum.
40///
41/// Executes the existing spatial and depth pooling kernels on the tensor's
42/// backend. Padding, dilation and ceil mode are applied independently per axis.
43/// Layout transformations and both pooling stages remain differentiable.
44pub 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
69/// Maximum volume pooling with original input positions in native I64 storage.
70///
71/// Indices flatten each input volume as `d * (H * W) + h * W + w`, independently
72/// for every batch/channel. Tie selection follows the
73/// underlying spatial kernel followed by the depth kernel. Invalid backend
74/// indices become `-1`; they are never used as out-of-bounds gather addresses.
75pub 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
118/// Average-pool native `[batch, channels, depth, height, width]` activations.
119///
120/// Dispatches native volume pooling, retaining backend padding and ceil-window
121/// semantics. Backends without a specialized volume kernel use native spatial
122/// and depth pooling; no host reduction is used.
123pub 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
135/// Adaptive average pooling to explicit `[depth, height, width]` extents.
136///
137/// Native adaptive pooling bins are retained, including overlapping bins and
138/// output extents larger than the input. Dispatches the backend's volume pooling
139/// operation; backends without a specialized kernel retain native separable pooling.
140pub 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
164/// Average pooling with explicit `(left, right)` padding.
165///
166/// When padding is excluded, the denominator counts original input elements,
167/// not the zeros materialized for asymmetric padding. Symmetric calls retain
168/// the original backend operation and arithmetic.
169pub 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
189/// Average pooling with explicit height and width `(before, after)` pairs.
190///
191/// Half-storage exclusion statistics use FP32; F64 inputs retain F64 statistics.
192/// The output retains the input storage, including its native gradient path.
193pub 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
213/// Average pooling with depth, height and width `(before, after)` pairs.
214///
215/// Coverage statistics broadcast across batch/channels and stay on the device.
216pub 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
236/// Maximum pooling with explicit `(left, right)` padding and native dilation.
237pub 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
254/// Maximum pooling with explicit height and width `(before, after)` pairs.
255pub 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
271/// Maximum pooling with explicit depth, height and width padding pairs.
272pub 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
327/// Maximum pooling with asymmetric length padding and original input positions.
328///
329/// Native symmetric indices retain their backend storage. Asymmetric indices
330/// are remapped on the device to I64 input coordinates; padding selections and
331/// native invalid indices are `-1`. Values retain the native gradient graph.
332pub 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
352/// Maximum pooling with asymmetric spatial padding and original `h * W + w` indices.
353///
354/// Indices exclude materialized padding, using `-1` for padding selections.
355pub 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
375/// Maximum volume pooling with asymmetric padding and original flattened positions.
376///
377/// Returned I64 indices address the unpadded depth/height/width volume. Invalid
378/// native indices and selected materialized padding are `-1`.
379pub 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}