use burn_core as burn;
use burn::config::Config;
use burn::module::{Content, DisplaySettings, Flag, Module, ModuleDisplay, Param};
use burn::tensor::{Distribution, Tensor};
#[derive(Config, Debug)]
pub struct GaussianNoiseConfig {
pub std: f64,
}
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct GaussianNoise {
pub std: f64,
pub training: Param<Flag>,
}
impl GaussianNoiseConfig {
pub fn init(&self) -> GaussianNoise {
if !self.std.is_finite() || self.std < 0.0 {
panic!(
"Standard deviation must be finite and non-negative, but got {}",
self.std
);
}
GaussianNoise {
std: self.std,
training: Param::from_bool(true),
}
}
}
impl GaussianNoise {
pub fn forward<const D: usize>(&self, input: Tensor<D>) -> Tensor<D> {
if self.training.is_enabled() && input.device().is_autodiff() && self.std != 0.0 {
let noise = Tensor::random(
input.shape(),
Distribution::Normal(0.0, self.std),
&input.device(),
);
input + noise
} else {
input
}
}
}
impl ModuleDisplay for GaussianNoise {
fn custom_settings(&self) -> Option<DisplaySettings> {
DisplaySettings::new()
.with_new_line_after_attribute(false)
.optional()
}
fn custom_content(&self, content: Content) -> Option<Content> {
let content = content.add("std", &self.std);
match self.training.is_enabled() {
true => content.optional(),
false => content.add("training", &self.training).optional(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::{Device, Shape};
#[cfg(feature = "std")]
#[test]
fn with_ad_backend_should_mark_input() {
let device = Device::default().autodiff();
let tensor = Tensor::<2>::ones(Shape::new([100, 100]), &device);
let noise = GaussianNoiseConfig::new(0.5).init();
let output = noise.forward(tensor.clone());
assert_ne!(tensor.to_data(), output.to_data());
}
#[test]
fn without_ad_backend_should_not_change_input() {
let tensor = Tensor::<2>::ones(Shape::new([100, 100]), &Default::default());
let noise = GaussianNoiseConfig::new(0.5).init();
let output = noise.forward(tensor.clone());
assert_eq!(tensor.to_data(), output.to_data());
}
#[cfg(feature = "std")]
#[test]
fn frozen_noise_on_a_training_device_passes_its_input_through() {
use burn::module::Module;
use burn::tensor::Device;
let device = Device::default().autodiff();
let tensor = Tensor::<2>::ones(Shape::new([100, 100]), &device);
let noise = GaussianNoiseConfig::new(0.5).init().freeze();
let output = noise.forward(tensor.clone());
assert_eq!(
tensor.to_data(),
output.to_data(),
"a frozen layer should not perturb a subtree the caller froze"
);
}
#[test]
#[should_panic(expected = "Standard deviation must be finite and non-negative")]
fn negative_std_should_panic() {
GaussianNoiseConfig { std: -0.5 }.init();
}
#[test]
#[should_panic(expected = "Standard deviation must be finite and non-negative")]
fn nan_std_should_panic() {
GaussianNoiseConfig::new(f64::NAN).init();
}
#[test]
#[should_panic(expected = "Standard deviation must be finite and non-negative")]
fn infinite_std_should_panic() {
GaussianNoiseConfig::new(f64::INFINITY).init();
}
#[test]
fn display() {
let config = GaussianNoiseConfig::new(0.5);
let layer = config.init();
assert_eq!(alloc::format!("{layer}"), "GaussianNoise {std: 0.5}");
}
#[test]
fn display_shows_a_frozen_layer() {
use burn::module::Module;
let layer = GaussianNoiseConfig::new(0.5).init().freeze();
assert_eq!(
alloc::format!("{layer}"),
"GaussianNoise {std: 0.5, training: disabled}"
);
}
}