1use crate::Float;
2
3pub(crate) fn sigmoid<E: Float>(i: E) -> E {
4 E::ONE / (E::ONE + i.neg().exp())
5}
6
7#[derive(Clone, Debug, Default)]
9pub enum Activation<E: Float> {
10 #[default]
14 Relu,
15 Sigmoid,
19 SiLU,
23 Tanh,
27 LeakyRelu(E),
29 Softplus,
31 Sine,
35 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 }
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}