use crate::{PaddingConfig3d, padding::dilated_kernel_size};
use ruda_model::{
config::Config,
module::{Content, DisplaySettings, Module, ModuleDisplay},
tensor::{Int, Tensor, backend::Backend,
module::{max_pool3d_padded, max_pool3d_with_indices_padded}},
};
#[derive(Config, Debug)]
pub struct MaxPool3dConfig {
pub kernel_size: [usize; 3],
#[config(default = "kernel_size")]
pub strides: [usize; 3],
#[config(default = "PaddingConfig3d::Valid")]
pub padding: PaddingConfig3d,
#[config(default = "[1, 1, 1]")]
pub dilation: [usize; 3],
#[config(default = "false")]
pub ceil_mode: bool,
}
#[derive(Module, Clone, Debug)]
#[module(custom_display)]
pub struct MaxPool3d {
pub stride: [usize; 3],
pub kernel_size: [usize; 3],
pub padding: PaddingConfig3d,
pub dilation: [usize; 3],
pub ceil_mode: bool,
}
impl MaxPool3dConfig {
pub fn init(&self) -> MaxPool3d {
MaxPool3d {
stride: self.strides,
kernel_size: self.kernel_size,
padding: self.padding.clone(),
dilation: self.dilation,
ceil_mode: self.ceil_mode,
}
}
}
impl MaxPool3d {
pub fn forward_with_indices<B: Backend>(&self, input: Tensor<B, 5>)
-> (Tensor<B, 5>, Tensor<B, 5, Int>) {
let [_, _, depth, height, width] = input.dims();
let effective = core::array::from_fn(|axis| {
dilated_kernel_size(self.kernel_size[axis], self.dilation[axis])
});
let padding = self.padding.calculate_padding_3d_pairs(
&[depth, height, width], &effective, &self.stride,
);
max_pool3d_with_indices_padded(input, self.kernel_size, self.stride, padding,
self.dilation, self.ceil_mode)
}
pub fn forward<B: Backend>(&self, input: Tensor<B, 5>) -> Tensor<B, 5> {
let [_, _, depth, height, width] = input.dims();
let effective = core::array::from_fn(|axis| {
dilated_kernel_size(self.kernel_size[axis], self.dilation[axis])
});
let pairs = self.padding.calculate_padding_3d_pairs(
&[depth, height, width], &effective, &self.stride,
);
max_pool3d_padded(input, self.kernel_size, self.stride, pairs, self.dilation, self.ceil_mode)
}
}
impl ModuleDisplay for MaxPool3d {
fn custom_settings(&self) -> Option<DisplaySettings> {
DisplaySettings::new().with_new_line_after_attribute(false).optional()
}
fn custom_content(&self, content: Content) -> Option<Content> {
content.add_debug_attribute("kernel_size", &self.kernel_size)
.add_debug_attribute("stride", &self.stride)
.add_debug_attribute("padding", &self.padding)
.add_debug_attribute("dilation", &self.dilation)
.add("ceil_mode", &self.ceil_mode).optional()
}
}