use crate::{Element, Value};
fn window_lanes<'tape, E: Element>(
input: Value<'tape, E>,
size: usize,
stride: usize,
) -> Value<'tape, E> {
let shape = input.shape();
assert_eq!(
shape.rank(),
4,
"pooling input must be rank 4 [batch, channels, height, width], got {shape}"
);
assert!(size > 0, "pooling windows must hold at least one element");
assert!(stride > 0, "pooling stride must be positive");
let windows = input.unfold(2, size, stride, 1).unfold(4, size, stride, 1);
let windows_shape = windows.shape();
let axes = windows_shape.axes();
windows
.permute([0, 1, 2, 4, 3, 5])
.reshape([axes[0], axes[1], axes[2], axes[4], size * size])
}
pub fn max_pool<'tape, E: Element>(
input: Value<'tape, E>,
size: usize,
stride: usize,
) -> Value<'tape, E> {
let lanes = window_lanes(input, size, stride);
let mut largest = lanes.narrow(4, 0, 1);
for lane in 1..size * size {
largest = largest.maximum(lanes.narrow(4, lane, 1));
}
largest.squeeze(4)
}
#[cfg(test)]
#[path = "tests/pooling_tests.rs"]
mod tests;