burn-nn 0.22.0

Neural network building blocks for the Burn deep learning framework
use burn_core as burn;

use burn::config::Config;
use burn::module::{Content, DisplaySettings, Flag, Module, ModuleDisplay, Param};
use burn::tensor::{Distribution, Tensor};

/// Configuration to create a [GaussianNoise](GaussianNoise) layer using the [init function](GaussianNoiseConfig::init).
#[derive(Config, Debug)]
pub struct GaussianNoiseConfig {
    /// Standard deviation of the normal noise distribution.
    pub std: f64,
}

/// Add pseudorandom Gaussian noise to an arbitrarily shaped tensor.
///
/// This is an effective regularization technique that also contributes to data augmentation.
/// Please keep in mind that the value of [std](GaussianNoise::std) should be chosen with care in order to avoid
/// distortion.
///
/// Should be created with [GaussianNoiseConfig].
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct GaussianNoise {
    /// Standard deviation of the normal noise distribution.
    pub std: f64,
    /// Whether to behave as during training. Cleared by
    /// [`freeze`](burn::module::Module::freeze) and matching
    /// [`freeze_group`](burn::module::Module::freeze_group) traversals.
    pub training: Param<Flag>,
}

impl GaussianNoiseConfig {
    /// Initialize a new [Gaussian noise](GaussianNoise) module.
    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 {
    /// Applies the forward pass on the input tensor.
    ///
    /// See [GaussianNoise](GaussianNoise) for more information.
    ///
    /// # Shapes
    ///
    /// - input: `[..., any]`
    /// - output: `[..., any]`
    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;
        // Frozen where partial finetuning leaves it: on the training device,
        // because that is where the rest of the graph is.
        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}"
        );
    }
}