use scirs2_core::ndarray::{Array, Dimension, ScalarOperand, Zip};
use scirs2_core::numeric::Float;
use scirs2_core::random::Rng;
use scirs2_core::Random;
use std::cell::RefCell;
use std::fmt::Debug;
use crate::error::Result;
use crate::regularizers::Regularizer;
#[derive(Debug)]
pub struct Dropout<A: Float + Debug> {
rate: A,
rng: RefCell<Random<scirs2_core::random::rngs::StdRng>>,
training: bool,
}
impl<A: Float + Debug + Send + Sync> Dropout<A> {
pub fn new<R: Rng>(rate: A, rng: &mut R) -> Self {
let rate = rate.max(A::zero()).min(A::one());
let mut seed_bytes = [0u8; 8];
rng.fill_bytes(&mut seed_bytes);
let seed = u64::from_ne_bytes(seed_bytes);
let rng = Random::seed(seed);
Self {
rate,
rng: RefCell::new(rng),
training: true,
}
}
pub fn rate(&self) -> A {
self.rate
}
pub fn set_rate(&mut self, rate: A) -> &mut Self {
self.rate = rate.max(A::zero()).min(A::one());
self
}
pub fn train(&mut self) -> &mut Self {
self.training = true;
self
}
pub fn eval(&mut self) -> &mut Self {
self.training = false;
self
}
pub fn is_training(&self) -> bool {
self.training
}
fn create_mask<D: Dimension>(&self, shape: D) -> Array<A, D> {
if !self.training || self.rate <= A::zero() {
return Array::ones(shape);
}
let keep_prob = A::one() - self.rate;
if keep_prob <= A::zero() {
return Array::zeros(shape);
}
let scale = A::one() / keep_prob;
let rate = self.rate.to_f64().unwrap_or(0.0);
let mut rng = self.rng.borrow_mut();
let mut mask = Array::zeros(shape);
for elem in mask.iter_mut() {
let rand_val: f64 = rng.gen_range(0.0..1.0);
if rand_val > rate {
*elem = scale;
}
}
mask
}
}
impl<A, D> Regularizer<A, D> for Dropout<A>
where
A: Float + ScalarOperand + Debug + Send + Sync,
D: Dimension<Pattern = D>,
{
fn apply(&self, _params: &Array<A, D>, gradients: &mut Array<A, D>) -> Result<A> {
if !self.training || self.rate <= A::zero() {
return Ok(A::zero());
}
let mask = self.create_mask(gradients.dim());
Zip::from(gradients).and(&mask).for_each(|grad, &mask_val| {
*grad = *grad * mask_val;
});
Ok(A::zero())
}
fn penalty(&self, _params: &Array<A, D>) -> Result<A> {
Ok(A::zero())
}
}
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::{Array1, ArrayD};
use scirs2_core::random::rngs::SmallRng;
use scirs2_core::random::SeedableRng;
fn make_dropout(rate: f64) -> Dropout<f64> {
let mut rng = SmallRng::from_seed([7u8; 32]);
Dropout::new(rate, &mut rng)
}
fn dyn_vec(values: &[f64]) -> ArrayD<f64> {
Array1::from_vec(values.to_vec()).into_dyn()
}
fn dyn_ones(len: usize) -> ArrayD<f64> {
Array1::from_elem(len, 1.0).into_dyn()
}
#[test]
fn eval_mode_leaves_gradients_untouched() {
let mut dropout = make_dropout(0.5);
dropout.eval();
let params = dyn_vec(&[1.0, 2.0, 3.0, 4.0]);
let original = dyn_vec(&[0.1, 0.2, 0.3, 0.4]);
let mut gradients = original.clone();
let penalty = dropout
.apply(¶ms, &mut gradients)
.expect("dropout apply failed");
assert_eq!(penalty, 0.0);
assert_eq!(gradients, original);
}
#[test]
fn training_mode_masks_gradients_not_params() {
let mut dropout = make_dropout(0.5);
dropout.train();
let params = dyn_ones(512);
let params_before = params.clone();
let mut gradients = dyn_ones(512);
dropout
.apply(¶ms, &mut gradients)
.expect("dropout apply failed");
assert_eq!(params, params_before);
assert!(gradients
.iter()
.all(|&g| g == 0.0 || (g - 2.0).abs() < 1e-12));
let dropped = gradients.iter().filter(|&&g| g == 0.0).count();
assert!(dropped > 0, "no gradient entries were dropped");
assert!(dropped < 512, "every gradient entry was dropped");
let sum: f64 = gradients.sum();
assert!((sum - 512.0).abs() < 160.0, "unexpected gradient sum {sum}");
}
#[test]
fn mask_is_redrawn_on_every_call() {
let mut dropout = make_dropout(0.5);
dropout.train();
let params = dyn_ones(256);
let mut first = dyn_ones(256);
let mut second = dyn_ones(256);
dropout
.apply(¶ms, &mut first)
.expect("dropout apply failed");
dropout
.apply(¶ms, &mut second)
.expect("dropout apply failed");
assert_ne!(first, second, "dropout mask appears to be cached");
}
#[test]
fn zero_rate_is_identity_and_full_rate_zeroes_everything() {
let params = dyn_vec(&[1.0, 2.0, 3.0]);
let original = dyn_vec(&[0.5, -1.5, 2.5]);
let mut none = make_dropout(0.0);
none.train();
let mut gradients = original.clone();
none.apply(¶ms, &mut gradients)
.expect("dropout apply failed");
assert_eq!(gradients, original);
let mut all = make_dropout(1.0);
all.train();
let mut gradients = original.clone();
all.apply(¶ms, &mut gradients)
.expect("dropout apply failed");
assert_eq!(gradients, dyn_vec(&[0.0, 0.0, 0.0]));
}
#[test]
fn penalty_is_always_zero() {
let dropout = make_dropout(0.5);
let params = dyn_vec(&[1.0, 2.0, 3.0]);
assert_eq!(
Regularizer::penalty(&dropout, ¶ms).expect("penalty failed"),
0.0
);
}
}