use burn::config::Config;
use burn::module::Module;
use burn::module::{Content, DisplaySettings, Flag, ModuleDisplay, Param};
use burn::tensor::{Distribution, Tensor};
use burn_core as burn;
use burn::tensor::activation::leaky_relu;
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct RRelu {
pub lower: f64,
pub upper: f64,
pub training: Param<Flag>,
}
#[derive(Config, Debug)]
pub struct RReluConfig {
#[config(default = "0.125")]
pub lower: f64,
#[config(default = "1.0 / 3.0")]
pub upper: f64,
}
impl RReluConfig {
pub fn init(&self) -> RRelu {
assert!(
self.lower <= self.upper,
"RRelu: lower bound ({}) must be <= upper bound ({})",
self.lower,
self.upper
);
RRelu {
lower: self.lower,
upper: self.upper,
training: Param::from_bool(true),
}
}
}
impl ModuleDisplay for RRelu {
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("lower", &self.lower).add("upper", &self.upper);
match self.training.is_enabled() {
true => content.optional(),
false => content.add("training", &self.training).optional(),
}
}
}
impl RRelu {
pub fn forward<const D: usize>(&self, input: Tensor<D>) -> Tensor<D> {
if !self.training.is_enabled() || !input.device().is_autodiff() {
return leaky_relu(input, (self.lower + self.upper) / 2.0);
}
let is_negative = input.clone().lower_scalar(0.0);
let slope = input.random_like(Distribution::Uniform(self.lower, self.upper));
let scaled = input.clone() * slope;
input.mask_where(is_negative, scaled)
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::TensorData;
use burn::tensor::Tolerance;
type FT = f32;
#[test]
fn eval_matches_leaky_relu_midpoint() {
let device = Default::default();
let model = RReluConfig::new().with_lower(0.1).with_upper(0.3).init();
let input = Tensor::<2>::from_data(
TensorData::from([[-2.0, -1.0, 0.0], [0.5, 1.0, 2.0]]),
&device,
);
let output = model.forward(input);
let expected = TensorData::from([[-0.4, -0.2, 0.0], [0.5, 1.0, 2.0]]);
output
.to_data()
.assert_approx_eq::<FT>(&expected, Tolerance::default());
}
#[cfg(feature = "std")]
#[test]
fn training_scales_negatives_randomly() {
use burn::tensor::Device;
let device = Device::default().autodiff();
let model = RReluConfig::new().init();
let input = Tensor::<2>::from_data(TensorData::from([[-1.0, -2.0], [-3.0, -4.0]]), &device);
let output = model.forward(input.clone());
assert_ne!(input.to_data(), output.to_data());
}
#[cfg(feature = "std")]
#[test]
fn frozen_rrelu_on_a_training_device_uses_the_fixed_slope() {
use burn::module::Module;
use burn::tensor::Device;
let device = Device::default().autodiff();
let model = RReluConfig::new().with_lower(0.1).with_upper(0.3).init();
let input = Tensor::<2>::from_data(TensorData::from([[-1.0, -2.0], [-3.0, -4.0]]), &device);
let output = model.freeze().forward(input.clone());
let expected = TensorData::from([[-0.2, -0.4], [-0.6, -0.8]]);
output
.to_data()
.assert_approx_eq::<FT>(&expected, Tolerance::default());
}
#[test]
fn display() {
let layer = RReluConfig::new().with_lower(0.1).with_upper(0.3).init();
assert_eq!(alloc::format!("{layer}"), "RRelu {lower: 0.1, upper: 0.3}");
}
#[test]
fn display_shows_a_frozen_layer() {
use burn::module::Module;
let layer = RReluConfig::new()
.with_lower(0.1)
.with_upper(0.3)
.init()
.freeze();
assert_eq!(
alloc::format!("{layer}"),
"RRelu {lower: 0.1, upper: 0.3, training: disabled}"
);
}
#[test]
#[should_panic = "must be <= upper bound"]
fn rejects_lower_above_upper() {
let _ = RReluConfig::new().with_lower(0.5).with_upper(0.2).init();
}
}