use alloc::format;
use burn::tensor::module::interpolate;
use burn_core as burn;
use burn::config::Config;
use burn::module::{Content, DisplaySettings, Module, ModuleDisplay};
use burn::tensor::Tensor;
use burn::tensor::ops::InterpolateOptions;
use super::InterpolateMode;
#[derive(Config, Debug)]
pub struct Interpolate1dConfig {
#[config(default = "None")]
pub output_size: Option<usize>,
#[config(default = "None")]
pub scale_factor: Option<f32>,
#[config(default = "InterpolateMode::Nearest")]
pub mode: InterpolateMode,
#[config(default = true)]
pub align_corners: bool,
}
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct Interpolate1d {
pub output_size: Option<usize>,
pub scale_factor: Option<f32>,
#[module(skip)]
pub mode: InterpolateMode,
pub align_corners: bool,
}
impl Interpolate1dConfig {
pub fn init(self) -> Interpolate1d {
Interpolate1d {
output_size: self.output_size,
scale_factor: self.scale_factor,
mode: self.mode,
align_corners: self.align_corners,
}
}
}
impl Interpolate1d {
pub fn forward(&self, input: Tensor<3>) -> Tensor<3> {
let mut options = InterpolateOptions::new(self.mode.clone().into())
.with_align_corners(self.align_corners);
options.output_size = self.output_size.map(|size| [1, size]);
options.scale_factor = self.scale_factor.map(|scale| [1.0, scale]);
let input = input.unsqueeze_dim(2);
let result = interpolate(input, options);
result.squeeze_dims(&[2])
}
}
impl ModuleDisplay for Interpolate1d {
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("output_size", &format!("{:?}", self.output_size))
.add("scale_factor", &self.scale_factor)
.optional()
}
}
#[cfg(test)]
mod tests {
use burn::tensor::Distribution;
use super::*;
#[test]
fn test_module() {
let input = Tensor::<3>::random(
[2, 3, 4],
Distribution::Uniform(0.0, 1.0),
&Default::default(),
);
let config = Interpolate1dConfig::new().with_output_size(Some(8));
let interpolate = config.init();
let output = interpolate.forward(input.clone());
assert_eq!(output.dims(), [2, 3, 8]);
let config = Interpolate1dConfig::new().with_scale_factor(Some(0.5));
let interpolate = config.init();
let output = interpolate.forward(input.clone());
assert_eq!(output.dims(), [2, 3, 2]);
let config = Interpolate1dConfig::new()
.with_output_size(Some(6))
.with_mode(InterpolateMode::Linear);
let interpolate = config.init();
let output = interpolate.forward(input);
assert_eq!(output.dims(), [2, 3, 6]);
}
#[test]
fn display() {
let config = Interpolate1dConfig::new().with_output_size(Some(20));
let layer = config.init();
assert_eq!(
alloc::format!("{layer}"),
"Interpolate1d {mode: Nearest, output_size: Some(20), \
scale_factor: None}"
);
}
}