Skip to main content

minidx_core/layers/
swish.rs

1use super::sigmoid;
2use crate::gradients::{ClassActivation, ClassWrapper, Gradients};
3use crate::Float;
4
5/// The swish activation function with learnable beta.
6#[derive(Clone, Debug)]
7pub struct Swish<E: Float, const I: usize> {
8    pub(crate) beta: ClassWrapper<[E; I], ClassActivation>,
9}
10
11impl<E: Float, const I: usize> Default for Swish<E, I> {
12    fn default() -> Self {
13        Self {
14            beta: ClassWrapper::<[E; I], ClassActivation>::wrap([E::ONE; I]),
15        }
16    }
17}
18
19impl<E: Float, const I: usize> Swish<E, I> {
20    #[inline]
21    fn forward(&self, input: &[E; I]) -> [E; I] {
22        let mut out: [E; I] = [E::default(); I];
23        for ((o, x), b) in out.iter_mut().zip(input.iter()).zip(self.beta.grad_iter()) {
24            *o = *x * sigmoid(*b * *x);
25        }
26        out
27    }
28
29    #[inline]
30    fn gradients_wrt_input(&self, input: &[E; I]) -> [E; I] {
31        let mut out: [E; I] = [E::default(); I];
32        for ((o, x), b) in out.iter_mut().zip(input.iter()).zip(self.beta.grad_iter()) {
33            let act = sigmoid(*b * *x);
34            *o = act * (E::ONE + (*b * *x) * E::ONE.sub(act));
35        }
36        out
37    }
38
39    #[inline]
40    fn gradients_wrt_beta(&self, input: &[E; I], output_gradients: &[E; I]) -> [E; I] {
41        let mut out: [E; I] = [E::default(); I];
42        for (((o, x), b), g) in out
43            .iter_mut()
44            .zip(input.iter())
45            .zip(self.beta.grad_iter())
46            .zip(output_gradients.iter())
47        {
48            let act = sigmoid(*b * *x);
49            *o = *g * (*x * *x) * act * E::ONE.sub(act);
50        }
51        out
52    }
53}
54
55impl<E: Float, const I: usize> crate::BaseModule for Swish<E, I> {}
56
57impl<E: Float, const I: usize> crate::Module<[E; I]> for Swish<E, I> {
58    type Output = [E; I];
59
60    fn forward(&self, x: &[E; I]) -> Result<Self::Output, crate::Error> {
61        Ok(Swish::forward(self, x))
62    }
63}
64
65impl<E: Float, const I: usize> crate::RevModule<[E; I]> for Swish<E, I> {
66    type SelfGrads = ClassWrapper<[E; I], ClassActivation>;
67
68    fn reverse(&self, inputs: &[E; I], grads_wrt_output: &[E; I]) -> ([E; I], Self::SelfGrads) {
69        let mut output_grads = self.gradients_wrt_input(inputs);
70        output_grads
71            .iter_mut()
72            .zip(grads_wrt_output)
73            .for_each(|(ga, go)| *ga *= *go);
74
75        (
76            output_grads,
77            Self::SelfGrads::wrap(self.gradients_wrt_beta(inputs, grads_wrt_output)),
78        )
79    }
80
81    fn apply(
82        &mut self,
83        applyer: &mut impl crate::optimizers::GradApplyer,
84        updates: Self::SelfGrads,
85    ) -> Result<(), crate::Error> {
86        applyer.apply(updates, &mut self.beta)
87    }
88}
89
90impl<E: Float, const I: usize> crate::LoadableModule for Swish<E, I> {
91    fn save(
92        &self,
93        path: String,
94        dict: &mut std::collections::HashMap<String, Vec<f64>>,
95    ) -> Result<(), crate::LoadSaveError> {
96        dict.insert(
97            path,
98            self.beta.grad_iter().map(|f| f.to_f64().unwrap()).collect(),
99        );
100        Ok(())
101    }
102
103    fn load(
104        &mut self,
105        path: String,
106        dict: &std::collections::HashMap<String, Vec<f64>>,
107    ) -> Result<(), crate::LoadSaveError> {
108        let params = dict.get(&path).ok_or(crate::LoadSaveError {
109            path: path.clone(),
110            err: "Parameters missing".into(),
111        })?;
112        if params.len() != I {
113            return Err(crate::LoadSaveError {
114                path,
115                err: format!(
116                    "Parameters have wrong size: got {}, want {}",
117                    params.len(),
118                    I
119                )
120                .into(),
121            });
122        }
123        for (a, b) in self.beta.grad_iter_mut().zip(params.into_iter()) {
124            *a = E::from_f64(*b).unwrap();
125        }
126        Ok(())
127    }
128}
129
130impl<E: Float, const I: usize> crate::ResetParams for Swish<E, I> {
131    fn rand_params<RNG: rand::Rng>(
132        &mut self,
133        rng: &mut RNG,
134        scale: f32,
135    ) -> Result<(), crate::Error> {
136        // Xavier/Glorot initialization vibes, but scaled down a bit
137        // and centered about 1.
138        let stddev = 1.0 / ((I * I) as f32 * 8.0).sqrt();
139        let normal = rand_distr::Normal::new(1.0, stddev).unwrap();
140
141        self.beta.grad_iter_mut().for_each(|b| {
142            let s: f32 = rng.sample::<f32, _>(normal) * scale;
143            *b = E::from_f32(s).unwrap();
144        });
145        Ok(())
146    }
147}
148
149impl<E: Float, const I: usize> crate::VisualizableUnit for Swish<E, I> {
150    const KIND: &'static str = "swish";
151    type Params = [[E; I]; 1];
152    fn params(&self) -> &Self::Params {
153        // SAFETY: An array of N is exactly the same as a unary array of the array of N
154        unsafe { std::mem::transmute(self.beta.raw_grads_ref()) }
155    }
156}
157
158#[cfg(test)]
159mod tests {
160    use super::*;
161
162    #[test]
163    fn test_default() {
164        let layer = Swish::<f32, 2>::default();
165        assert_eq!(layer.beta.raw_grads(), [1.0, 1.0],);
166    }
167
168    #[test]
169    fn test_forward() {
170        let mut layer = Swish::<f32, 4>::default();
171        layer.beta.raw_grads_mut()[2] = 355.0;
172        let out = layer.forward(&[10.0, 0.0, 1.0, 1.0]);
173        assert!(out[0] > 9.99 && out[0] < 10.0);
174        assert_eq!(out[1], 0.0);
175        assert_eq!(out[2], 1.0);
176        assert!(out[3] > 0.72 && out[3] < 0.75);
177    }
178
179    #[test]
180    fn test_gradients_wrt_input() {
181        let mut layer = Swish::<f32, 1>::default();
182        let out = layer.gradients_wrt_input(&[2.0]);
183        assert!(out[0] > 1.08 && out[0] < 1.091);
184
185        layer.beta.raw_grads_mut()[0] = 10.0;
186        let out = layer.gradients_wrt_input(&[2.0]);
187        assert!(out[0] > 0.9999 && out[0] < 1.0001);
188    }
189
190    #[test]
191    fn test_gradients_wrt_beta() {
192        let mut layer = Swish::<f32, 1>::default();
193        let out = layer.gradients_wrt_beta(&[2.0], &[1.0]);
194        assert!(out[0] > 0.409 && out[0] < 0.421);
195
196        layer.beta.raw_grads_mut()[0] = 10.0;
197        let out = layer.gradients_wrt_beta(&[-2.0], &[1.0]);
198        assert!(out[0] > 8.0e-9 && out[0] < 8.26e-9);
199    }
200}