use ruda_model::{
config::Config,
module::{Content, DisplaySettings, Module, ModuleDisplay},
tensor::{Tensor, backend::Backend, module::interpolate3d, ops::InterpolateOptions},
};
use super::InterpolateMode;
#[derive(Config, Debug)]
pub struct Interpolate3dConfig {
#[config(default = "None")]
pub output_size: Option<[usize; 3]>,
#[config(default = "None")]
pub scale_factor: Option<[f32; 3]>,
#[config(default = "InterpolateMode::Nearest")]
pub mode: InterpolateMode,
#[config(default = true)]
pub align_corners: bool,
}
#[derive(Module, Clone, Debug)]
#[module(custom_display)]
pub struct Interpolate3d {
pub output_size: Option<[usize; 3]>,
pub scale_factor: Option<[f32; 3]>,
pub mode: InterpolateMode,
pub align_corners: bool,
}
impl Interpolate3dConfig {
pub fn init(self) -> Interpolate3d {
Interpolate3d {
output_size: self.output_size,
scale_factor: self.scale_factor,
mode: self.mode,
align_corners: self.align_corners,
}
}
}
impl Interpolate3d {
pub fn forward_with_output_size<B: Backend>(
&self,
input: Tensor<B, 5>,
output_size: [usize; 3],
) -> Tensor<B, 5> {
interpolate3d(input, output_size,
InterpolateOptions::new(self.mode.clone().into()).with_align_corners(self.align_corners))
}
pub fn forward<B: Backend>(&self, input: Tensor<B, 5>) -> Tensor<B, 5> {
let output = if let Some(size) = self.output_size {
assert!(size.iter().all(|size| *size > 0), "interpolation output extents must be non-zero");
size
} else {
let factors = self.scale_factor.expect("Either output_size or scale_factor must be provided");
let [_, _, depth, height, width] = input.dims();
let sizes = [depth, height, width];
core::array::from_fn(|axis| {
let factor = factors[axis];
assert!(factor.is_finite() && factor > 0.0, "interpolation scale factors must be finite and positive");
let size = (sizes[axis] as f64) * (factor as f64);
assert!(size >= 1.0 && size < usize::MAX as f64,
"interpolation scale factor produces an empty or overflowing output extent");
size as usize
})
};
interpolate3d(input, output,
InterpolateOptions::new(self.mode.clone().into()).with_align_corners(self.align_corners))
}
}
impl ModuleDisplay for Interpolate3d {
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("mode", &self.mode)
.add_debug_attribute("output_size", &self.output_size)
.add_debug_attribute("scale_factor", &self.scale_factor)
.add("align_corners", &self.align_corners).optional()
}
}