use burn_core as burn;
use burn::config::Config;
use burn::module::{Content, DisplaySettings, Module, ModuleDisplay};
use burn::tensor::Tensor;
use burn::tensor::module::fold4d;
use burn::tensor::ops::UnfoldOptions;
#[derive(Config, Debug)]
pub struct Fold4dConfig {
pub output_size: [usize; 2],
pub kernel_size: [usize; 2],
#[config(default = "[1, 1]")]
pub stride: [usize; 2],
#[config(default = "[1, 1]")]
pub dilation: [usize; 2],
#[config(default = "[0, 0]")]
pub padding: [usize; 2],
}
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct Fold4d {
pub output_size: [usize; 2],
pub kernel_size: [usize; 2],
pub stride: [usize; 2],
pub dilation: [usize; 2],
pub padding: [usize; 2],
}
impl ModuleDisplay for Fold4d {
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("output_size", &alloc::format!("{:?}", self.output_size))
.add("kernel_size", &alloc::format!("{:?}", self.kernel_size))
.add("stride", &alloc::format!("{:?}", self.stride))
.add("dilation", &alloc::format!("{:?}", self.dilation))
.add("padding", &alloc::format!("{:?}", self.padding))
.optional()
}
}
impl Fold4dConfig {
pub fn init(&self) -> Fold4d {
Fold4d {
output_size: self.output_size,
kernel_size: self.kernel_size,
stride: self.stride,
dilation: self.dilation,
padding: self.padding,
}
}
}
impl Fold4d {
pub fn forward(&self, input: Tensor<3>) -> Tensor<4> {
fold4d(
input,
self.output_size,
self.kernel_size,
UnfoldOptions::new(self.stride, self.padding, self.dilation),
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::module::unfold4d;
use burn::tensor::{TensorData, Tolerance};
type FT = f32;
#[test]
fn display() {
let config = Fold4dConfig::new([4, 4], [2, 2]);
let fold = config.init();
assert_eq!(
alloc::format!("{fold}"),
"Fold4d {output_size: [4, 4], kernel_size: [2, 2], stride: [1, 1], dilation: [1, 1], padding: [0, 0]}"
);
}
#[test]
fn fold4d_known_values() {
let device = Default::default();
let cols = Tensor::<3>::from_data(
TensorData::from([[
[1.0, 2.0, 3.0, 4.0],
[5.0, 6.0, 7.0, 8.0],
[9.0, 10.0, 11.0, 12.0],
[13.0, 14.0, 15.0, 16.0],
]]),
&device,
);
let output = Fold4dConfig::new([3, 3], [2, 2]).init().forward(cols);
let expected =
TensorData::from([[[[1.0, 7.0, 6.0], [12.0, 34.0, 22.0], [11.0, 27.0, 16.0]]]]);
output
.to_data()
.assert_approx_eq::<FT>(&expected, Tolerance::default());
}
#[test]
fn fold_of_unfold_matches_reference() {
let device = Default::default();
let x = Tensor::<4>::from_data(
TensorData::from([[[
[1.0, 2.0, 3.0, 4.0],
[5.0, 6.0, 7.0, 8.0],
[9.0, 10.0, 11.0, 12.0],
[13.0, 14.0, 15.0, 16.0],
]]]),
&device,
);
let cols = unfold4d(x, [2, 2], UnfoldOptions::new([1, 1], [0, 0], [1, 1]));
let folded = Fold4dConfig::new([4, 4], [2, 2]).init().forward(cols);
let expected = TensorData::from([[[
[1.0, 4.0, 6.0, 4.0],
[10.0, 24.0, 28.0, 16.0],
[18.0, 40.0, 44.0, 24.0],
[13.0, 28.0, 30.0, 16.0],
]]]);
folded
.to_data()
.assert_approx_eq::<FT>(&expected, Tolerance::default());
}
}