Skip to main content

minidx_core/layers/
activation.rs

1use crate::Float;
2
3pub(crate) fn sigmoid<E: Float>(i: E) -> E {
4    E::ONE / (E::ONE + i.neg().exp())
5}
6
7/// An element-wise activation function with no trainable parameters.
8#[derive(Clone, Debug, Default)]
9pub enum Activation<E: Float> {
10    /// [Rectified Linear Unit (ReLU)](https://en.wikipedia.org/wiki/Rectifier_(neural_networks)). `max(0, t)`
11    ///
12    /// The derivative is the [Heaviside](https://en.wikipedia.org/wiki/Heaviside_step_function) function.
13    #[default]
14    Relu,
15    /// [Sigmoid](https://en.wikipedia.org/wiki/Sigmoid_function). `1 / (1 + exp(-t))`.
16    ///
17    /// The derivative is `sigmoid(t) * (1.0 - sigmoid(t))`.
18    Sigmoid,
19    /// [SiLU / Swish1](https://en.wikipedia.org/wiki/Swish_function). `t / (1 + exp(-t))`.
20    ///
21    /// The derivative is `sigmoid(t) * (1.0 + t * (1.0 - sigmoid(t)))`.
22    SiLU,
23    /// [Tanh](https://en.wikipedia.org/wiki/Hyperbolic_functions). `(e^x - e^-x) / (e^x + e^-x)`.
24    ///
25    /// The derivative is `1 - (tanh(t)^2)`.
26    Tanh,
27    /// [Leaky ReLu](https://en.wikipedia.org/wiki/Rectifier_(neural_networks)#Piecewise-linear_variants). `if t > 0 { t } else { a * t }`
28    LeakyRelu(E),
29    /// [Softplus](https://en.wikipedia.org/wiki/Softplus). `ln(1 + e^x)`
30    Softplus,
31    /// Sine function.
32    ///
33    /// The derivative is `cos(t)`.
34    Sine,
35    /// Cosine function.
36    ///
37    /// The derivative is `-sin(t)`.
38    Cosine,
39}
40
41impl<E: Float> Activation<E> {
42    #[inline]
43    fn forward<const I: usize>(&self, input: &[E; I]) -> [E; I] {
44        let mut out: [E; I] = [E::default(); I];
45        for (o, i) in out.iter_mut().zip(input.iter()) {
46            *o = match self {
47                Activation::Sigmoid => sigmoid(*i),
48                Activation::SiLU => *i * sigmoid(*i),
49                Activation::Tanh => i.tanh(),
50                Activation::Relu => E::default().max(*i),
51                Activation::LeakyRelu(a) => {
52                    if i < &E::default() {
53                        *a * *i
54                    } else {
55                        *i
56                    }
57                }
58                Activation::Softplus => (i.exp() + E::ONE).ln(),
59                Activation::Sine => i.sin(),
60                Activation::Cosine => i.cos(),
61            };
62        }
63        out
64    }
65
66    #[inline]
67    fn backward<const I: usize>(&self, input: &[E; I]) -> [E; I] {
68        let mut out: [E; I] = [E::default(); I];
69        for (o, i) in out.iter_mut().zip(input.iter()) {
70            *o = match self {
71                Activation::Sigmoid => {
72                    let sig = sigmoid(*i);
73                    sig * E::ONE.sub(sig)
74                    // TODO: Do we need to compute sigmoid, can we just use i?
75                    // Thats what dfdx does: https://github.com/coreylowman/dfdx/blob/main/dfdx-core/src/tensor_ops/sigmoid/cpu_kernel.rs#L12
76                }
77                Activation::SiLU => {
78                    let sig = sigmoid(*i);
79                    sig * (E::ONE + *i * E::ONE.sub(sig))
80                }
81                Activation::Tanh => {
82                    let tanh = i.tanh();
83                    E::ONE.sub(tanh * tanh)
84                }
85                Activation::Relu => {
86                    if i > &E::default() {
87                        E::ONE
88                    } else {
89                        E::default()
90                    }
91                }
92                Activation::LeakyRelu(a) => {
93                    if i < &E::default() {
94                        *a
95                    } else {
96                        E::ONE
97                    }
98                }
99                Activation::Softplus => sigmoid(*i),
100                Activation::Sine => (*i).cos(),
101                Activation::Cosine => -(*i).sin(),
102            };
103        }
104        out
105    }
106}
107
108impl<E: Float> crate::BaseModule for Activation<E> {}
109
110impl<E: Float, const I: usize> crate::Module<[E; I]> for Activation<E> {
111    type Output = [E; I];
112
113    fn forward(&self, x: &[E; I]) -> Result<Self::Output, crate::Error> {
114        Ok(Activation::forward(self, x))
115    }
116}
117
118impl<E: Float, const I: usize> crate::RevModule<[E; I]> for Activation<E> {
119    type SelfGrads = ();
120
121    fn reverse(&self, inputs: &[E; I], grads_wrt_output: &[E; I]) -> ([E; I], Self::SelfGrads) {
122        let mut grads = self.backward(inputs);
123        grads
124            .iter_mut()
125            .zip(grads_wrt_output)
126            .for_each(|(ga, go)| *ga *= *go);
127
128        (grads, ())
129    }
130
131    fn apply(
132        &mut self,
133        _applyer: &mut impl crate::optimizers::GradApplyer,
134        _updates: Self::SelfGrads,
135    ) -> Result<(), crate::Error> {
136        Ok(())
137    }
138}
139
140impl<E: Float> crate::ResetParams for Activation<E> {
141    fn rand_params<RNG: rand::Rng>(
142        &mut self,
143        _rng: &mut RNG,
144        _scale: f32,
145    ) -> Result<(), crate::Error> {
146        Ok(())
147    }
148}
149
150impl<E: Float> crate::LoadableModule for Activation<E> {
151    fn save(
152        &self,
153        _path: String,
154        _dict: &mut std::collections::HashMap<String, Vec<f64>>,
155    ) -> Result<(), crate::LoadSaveError> {
156        Ok(())
157    }
158
159    fn load(
160        &mut self,
161        _path: String,
162        _dict: &std::collections::HashMap<String, Vec<f64>>,
163    ) -> Result<(), crate::LoadSaveError> {
164        Ok(())
165    }
166}
167
168impl<E: Float> crate::VisualizableUnit for Activation<E> {
169    const KIND: &'static str = "activation";
170    type Params = ();
171    fn params(&self) -> &Self::Params {
172        &()
173    }
174}
175
176#[cfg(test)]
177mod tests {
178    use super::*;
179
180    #[test]
181    fn test_sigmoid() {
182        let layer = Activation::Sigmoid;
183        let out = layer.forward(&[10.0, -10.0]);
184        assert!(out[0] > 0.99994);
185        assert!(out[0] < 0.99996);
186        assert!(out[1] > 1.0e-6);
187        assert!(out[1] < 1.0e-4);
188    }
189
190    #[test]
191    fn test_silu() {
192        let layer = Activation::SiLU;
193        let out = layer.forward(&[10.0, -10.0, 2.0]);
194        assert!(out[0] > 9.9, "val is {}", out[0]);
195        assert!(out[0] < 10.0, "val is {}", out[0]);
196        assert!(out[1] > -1.0e-2, "val is {}", out[1]);
197        assert!(out[1] < -1.0e-5, "val is {}", out[1]);
198        assert!(out[2] > 1.75, "val is {}", out[2]);
199        assert!(out[2] < 1.77, "val is {}", out[2]);
200
201        let back = layer.backward(&[2.0]);
202        assert!(back[0] > 1.08, "val is {}", back[0]);
203        assert!(back[0] < 1.10, "val is {}", back[0]);
204    }
205
206    #[test]
207    fn test_tanh() {
208        let layer = Activation::Tanh;
209        let out = layer.forward(&[10.0, -10.0, 0.0]);
210        assert!(out[0] > 0.9, "val is {}", out[0]);
211        assert!(out[0] < 1.01, "val is {}", out[0]);
212        assert!(out[1] > -1.1, "val is {}", out[1]);
213        assert!(out[1] < -0.9, "val is {}", out[1]);
214        assert!(out[2] == 0.0, "val is {}", out[2]);
215
216        let back = layer.backward(&[0.0, -3.0]);
217        assert!(back[0] > 0.99, "val is {}", back[0]);
218        assert!(back[0] < 1.01, "val is {}", back[0]);
219        assert!(back[1] < 0.01, "val is {}", back[0]);
220    }
221
222    #[test]
223    fn test_relu() {
224        let layer = Activation::Relu;
225        let out = layer.forward(&[10.0, 0.0, -0.0001]);
226        assert_eq!(out[0], 10.0);
227        assert_eq!(out[1], 0.0);
228        assert_eq!(out[2], 0.0);
229    }
230
231    #[test]
232    fn test_leaky_relu() {
233        let layer = Activation::LeakyRelu(0.01);
234        let out = layer.forward(&[10.0, 0.0, -0.1]);
235        assert_eq!(out[0], 10.0);
236        assert_eq!(out[1], 0.0);
237        assert_eq!(out[2], -0.001);
238    }
239
240    #[test]
241    fn test_softplus() {
242        let layer = Activation::Softplus;
243        let out = layer.forward(&[1.0]);
244        assert!(out[0] > 1.31);
245        assert!(out[0] < 1.33);
246
247        let back = layer.backward(&[1.0]);
248        assert!(back[0] > 0.72, "val is {}", back[0]);
249        assert!(back[0] < 0.74, "val is {}", back[0]);
250    }
251}