use super::{BasicOps, DType, Int, Tensor, TensorPrimitive, backend::Backend};
use super::ops::PadMode;
use super::module::{
avg_pool1d, avg_pool2d,
max_pool1d, max_pool2d, max_pool1d_with_indices, max_pool2d_with_indices,
};
pub(super) fn volume_planes<B: Backend, K: BasicOps<B>>(input: Tensor<B, 5, K>) -> Tensor<B, 4, K> {
let [batch, channels, depth, height, width] = input.dims();
let planes = batch.checked_mul(depth).expect("pooling plane count overflow");
input.permute([0, 2, 1, 3, 4]).reshape([planes, channels, height, width])
}
pub(super) fn plane_depth_lines<B: Backend, K: BasicOps<B>>(
planes: Tensor<B, 4, K>,
batch: usize,
depth: usize,
) -> Tensor<B, 3, K> {
let [_, channels, height, width] = planes.dims();
let lines = batch.checked_mul(channels)
.and_then(|count| count.checked_mul(height))
.and_then(|count| count.checked_mul(width))
.expect("pooling depth line count overflow");
planes.reshape([batch, depth, channels, height, width])
.permute([0, 2, 3, 4, 1]).reshape([lines, 1, depth])
}
pub(super) fn depth_lines_volume<B: Backend, K: BasicOps<B>>(
lines: Tensor<B, 3, K>,
batch: usize,
channels: usize,
height: usize,
width: usize,
) -> Tensor<B, 5, K> {
let [_, _, depth] = lines.dims();
lines.reshape([batch, channels, height, width, depth]).permute([0, 1, 4, 2, 3])
}
pub fn max_pool3d<B: Backend>(
input: Tensor<B, 5>,
kernel_size: [usize; 3],
stride: [usize; 3],
padding: [usize; 3],
dilation: [usize; 3],
ceil_mode: bool,
) -> Tensor<B, 5> {
let [batch, channels, depth, _, _] = input.dims();
let planes = max_pool2d(
volume_planes(input),
[kernel_size[1], kernel_size[2]],
[stride[1], stride[2]],
[padding[1], padding[2]],
[dilation[1], dilation[2]],
ceil_mode,
);
let [_, _, height, width] = planes.dims();
let lines = max_pool1d(
plane_depth_lines(planes, batch, depth),
kernel_size[0], stride[0], padding[0], dilation[0], ceil_mode,
);
depth_lines_volume(lines, batch, channels, height, width)
}
pub fn max_pool3d_with_indices<B: Backend>(
input: Tensor<B, 5>,
kernel_size: [usize; 3],
stride: [usize; 3],
padding: [usize; 3],
dilation: [usize; 3],
ceil_mode: bool,
) -> (Tensor<B, 5>, Tensor<B, 5, Int>) {
let [batch, channels, depth, input_height, input_width] = input.dims();
assert!(depth > 0 && input_height > 0 && input_width > 0,
"volume pooling indices require non-empty spatial axes");
let area = input_height.checked_mul(input_width).expect("pooling plane area overflow");
let volume = depth.checked_mul(area).expect("pooling volume size overflow");
assert!(volume <= i64::MAX as usize, "volume pooling indices exceed I64");
let (planes, plane_indices) = max_pool2d_with_indices(
volume_planes(input),
[kernel_size[1], kernel_size[2]],
[stride[1], stride[2]],
[padding[1], padding[2]],
[dilation[1], dilation[2]],
ceil_mode,
);
let [_, _, height, width] = planes.dims();
let (lines, depth_indices) = max_pool1d_with_indices(
plane_depth_lines(planes, batch, depth),
kernel_size[0], stride[0], padding[0], dilation[0], ceil_mode,
);
let depth_indices = depth_indices.cast(DType::I64);
let safe_depth_indices = depth_indices.clone().clamp(0, depth as i64 - 1);
let spatial_indices = plane_depth_lines(plane_indices.cast(DType::I64), batch, depth)
.gather(2, safe_depth_indices.clone());
let invalid = depth_indices.clone().lower_elem(0)
.bool_or(depth_indices.clone().greater_equal_elem(depth as i64))
.bool_or(spatial_indices.clone().lower_elem(0))
.bool_or(spatial_indices.clone().greater_equal_elem(area as i64));
let indices = (safe_depth_indices.mul_scalar(area as i64)
+ spatial_indices.clamp(0, area as i64 - 1)).mask_fill(invalid, -1);
(
depth_lines_volume(lines, batch, channels, height, width),
depth_lines_volume(indices, batch, channels, height, width),
)
}
pub fn avg_pool3d<B: Backend>(
input: Tensor<B, 5>,
kernel_size: [usize; 3],
stride: [usize; 3],
padding: [usize; 3],
count_include_pad: bool,
ceil_mode: bool,
) -> Tensor<B, 5> {
let [batch, channels, depth, _, _] = input.dims();
let planes = avg_pool2d(
volume_planes(input),
[kernel_size[1], kernel_size[2]],
[stride[1], stride[2]],
[padding[1], padding[2]],
count_include_pad,
ceil_mode,
);
let [_, _, height, width] = planes.dims();
let lines = avg_pool1d(
plane_depth_lines(planes, batch, depth),
kernel_size[0], stride[0], padding[0], count_include_pad, ceil_mode,
);
depth_lines_volume(lines, batch, channels, height, width)
}
pub fn adaptive_avg_pool3d<B: Backend>(
input: Tensor<B, 5>,
output_size: [usize; 3],
) -> Tensor<B, 5> {
Tensor::new(TensorPrimitive::Float(B::adaptive_avg_pool3d(input.primitive.tensor(), output_size)))
}
fn average_excluding_explicit_padding<B: Backend, const D: usize, const N: usize>(
input: Tensor<B, D>,
padding: [(usize, usize); N],
pool: impl Fn(Tensor<B, D>) -> Tensor<B, D>,
) -> Tensor<B, D> {
let storage = input.dtype();
let compute = if storage == DType::F64 { DType::F64 } else { DType::F32 };
let mut visible_shape = input.dims();
visible_shape[0] = 1;
visible_shape[1] = 1;
let visible = Tensor::<B, D>::ones(visible_shape, (&input.device(), compute))
.pad(padding, PadMode::Constant(0.0));
let values = pool(input.cast(compute).pad(padding, PadMode::Constant(0.0)));
let coverage = pool(visible);
(values / coverage).cast(storage)
}
pub fn avg_pool1d_padded<B: Backend>(
input: Tensor<B, 3>,
kernel_size: usize,
stride: usize,
padding: [(usize, usize); 1],
count_include_pad: bool,
ceil_mode: bool,
) -> Tensor<B, 3> {
let [(left, right)] = padding;
if left == right {
return avg_pool1d(input, kernel_size, stride, left, count_include_pad, ceil_mode);
}
if count_include_pad {
return avg_pool1d(input.pad(padding, PadMode::Constant(0.0)),
kernel_size, stride, 0, true, ceil_mode);
}
average_excluding_explicit_padding(input, padding,
|input| avg_pool1d(input, kernel_size, stride, 0, true, ceil_mode))
}
pub fn avg_pool2d_padded<B: Backend>(
input: Tensor<B, 4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [(usize, usize); 2],
count_include_pad: bool,
ceil_mode: bool,
) -> Tensor<B, 4> {
if padding.iter().all(|(start, end)| start == end) {
return avg_pool2d(input, kernel_size, stride, padding.map(|(start, _)| start),
count_include_pad, ceil_mode);
}
if count_include_pad {
return avg_pool2d(input.pad(padding, PadMode::Constant(0.0)),
kernel_size, stride, [0; 2], true, ceil_mode);
}
average_excluding_explicit_padding(input, padding,
|input| avg_pool2d(input, kernel_size, stride, [0; 2], true, ceil_mode))
}
pub fn avg_pool3d_padded<B: Backend>(
input: Tensor<B, 5>,
kernel_size: [usize; 3],
stride: [usize; 3],
padding: [(usize, usize); 3],
count_include_pad: bool,
ceil_mode: bool,
) -> Tensor<B, 5> {
if padding.iter().all(|(start, end)| start == end) {
return avg_pool3d(input, kernel_size, stride, padding.map(|(start, _)| start),
count_include_pad, ceil_mode);
}
if count_include_pad {
return avg_pool3d(input.pad(padding, PadMode::Constant(0.0)),
kernel_size, stride, [0; 3], true, ceil_mode);
}
average_excluding_explicit_padding(input, padding,
|input| avg_pool3d(input, kernel_size, stride, [0; 3], true, ceil_mode))
}
pub fn max_pool1d_padded<B: Backend>(
input: Tensor<B, 3>,
kernel_size: usize,
stride: usize,
padding: [(usize, usize); 1],
dilation: usize,
ceil_mode: bool,
) -> Tensor<B, 3> {
let [(left, right)] = padding;
let (input, padding) = if left == right {
(input, left)
} else {
(input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)), 0)
};
max_pool1d(input, kernel_size, stride, padding, dilation, ceil_mode)
}
pub fn max_pool2d_padded<B: Backend>(
input: Tensor<B, 4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [(usize, usize); 2],
dilation: [usize; 2],
ceil_mode: bool,
) -> Tensor<B, 4> {
let (input, padding) = if padding.iter().all(|(start, end)| start == end) {
(input, padding.map(|(start, _)| start))
} else {
(input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)), [0; 2])
};
max_pool2d(input, kernel_size, stride, padding, dilation, ceil_mode)
}
pub fn max_pool3d_padded<B: Backend>(
input: Tensor<B, 5>,
kernel_size: [usize; 3],
stride: [usize; 3],
padding: [(usize, usize); 3],
dilation: [usize; 3],
ceil_mode: bool,
) -> Tensor<B, 5> {
let (input, padding) = if padding.iter().all(|(start, end)| start == end) {
(input, padding.map(|(start, _)| start))
} else {
(input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)), [0; 3])
};
max_pool3d(input, kernel_size, stride, padding, dilation, ceil_mode)
}
fn unpad_pool_indices<B: Backend, const D: usize, const N: usize>(
indices: Tensor<B, D, Int>,
input_size: [usize; N],
padding: [(usize, usize); N],
) -> Tensor<B, D, Int> {
let padded_size: [usize; N] = core::array::from_fn(|axis| {
input_size[axis].checked_add(padding[axis].0)
.and_then(|size| size.checked_add(padding[axis].1))
.expect("padded pooling index extent overflow")
});
let padded_volume = padded_size.iter().try_fold(1usize,
|size, axis| size.checked_mul(*axis)).expect("padded pooling index volume overflow");
let input_volume = input_size.iter().try_fold(1usize,
|size, axis| size.checked_mul(*axis)).expect("pooling input index volume overflow");
assert!(padded_volume > 0 && padded_volume <= i64::MAX as usize
&& input_volume <= i64::MAX as usize, "pooling indices cannot be represented in I64");
let indices = indices.cast(DType::I64);
if input_volume == 0 {
return indices.zeros_like().sub_scalar(1);
}
let mut invalid = indices.clone().lower_elem(0)
.bool_or(indices.clone().greater_equal_elem(padded_volume as i64));
let mut remaining = indices.clone().clamp(0, padded_volume as i64 - 1);
let mut unpadded = indices.zeros_like();
let mut input_stride = 1usize;
for axis in (0..N).rev() {
let coordinate = remaining.clone().remainder_scalar(padded_size[axis] as i64)
.sub_scalar(padding[axis].0 as i64);
remaining = remaining.div_scalar(padded_size[axis] as i64);
invalid = invalid.bool_or(coordinate.clone().lower_elem(0))
.bool_or(coordinate.clone().greater_equal_elem(input_size[axis] as i64));
unpadded = unpadded + coordinate.clamp(0, input_size[axis].saturating_sub(1) as i64)
.mul_scalar(input_stride as i64);
input_stride = input_stride.checked_mul(input_size[axis])
.expect("pooling input index stride overflow");
}
unpadded.mask_fill(invalid, -1)
}
pub fn max_pool1d_with_indices_padded<B: Backend>(
input: Tensor<B, 3>,
kernel_size: usize,
stride: usize,
padding: [(usize, usize); 1],
dilation: usize,
ceil_mode: bool,
) -> (Tensor<B, 3>, Tensor<B, 3, Int>) {
let [(left, right)] = padding;
if left == right {
return max_pool1d_with_indices(input, kernel_size, stride, left, dilation, ceil_mode);
}
let [_, _, length] = input.dims();
let (values, indices) = max_pool1d_with_indices(
input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)),
kernel_size, stride, 0, dilation, ceil_mode,
);
(values, unpad_pool_indices(indices, [length], padding))
}
pub fn max_pool2d_with_indices_padded<B: Backend>(
input: Tensor<B, 4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [(usize, usize); 2],
dilation: [usize; 2],
ceil_mode: bool,
) -> (Tensor<B, 4>, Tensor<B, 4, Int>) {
if padding.iter().all(|(start, end)| start == end) {
return max_pool2d_with_indices(input, kernel_size, stride,
padding.map(|(start, _)| start), dilation, ceil_mode);
}
let [_, _, height, width] = input.dims();
let (values, indices) = max_pool2d_with_indices(
input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)),
kernel_size, stride, [0; 2], dilation, ceil_mode,
);
(values, unpad_pool_indices(indices, [height, width], padding))
}
pub fn max_pool3d_with_indices_padded<B: Backend>(
input: Tensor<B, 5>,
kernel_size: [usize; 3],
stride: [usize; 3],
padding: [(usize, usize); 3],
dilation: [usize; 3],
ceil_mode: bool,
) -> (Tensor<B, 5>, Tensor<B, 5, Int>) {
if padding.iter().all(|(start, end)| start == end) {
return max_pool3d_with_indices(input, kernel_size, stride,
padding.map(|(start, _)| start), dilation, ceil_mode);
}
let [_, _, depth, height, width] = input.dims();
let (values, indices) = max_pool3d_with_indices(
input.pad(padding, PadMode::Constant(f32::NEG_INFINITY)),
kernel_size, stride, [0; 3], dilation, ceil_mode,
);
(values, unpad_pool_indices(indices, [depth, height, width], padding))
}