Skip to main content

only_brain/
activation_functions.rs

1use serde::{Deserialize, Serialize};
2
3/// The function a layer applies to each neuron's weighted sum.
4///
5/// Every layer of a [`crate::NeuralNetwork`] has its own, so hidden layers and the
6/// output layer can differ: a common choice is [`ReLU`](Self::ReLU) or
7/// [`Tanh`](Self::Tanh) for hidden layers, with [`Tanh`](Self::Tanh) for outputs in
8/// `[-1, 1]`, [`Sigmoid`](Self::Sigmoid) for outputs in `[0, 1]`, or
9/// [`Identity`](Self::Identity) for unbounded outputs.
10#[derive(Clone, Copy, PartialEq, Eq, Hash, Debug, Default, Serialize, Deserialize)]
11#[non_exhaustive]
12pub enum ActivationFunction {
13    /// `1 / (1 + e^-x)`, in `(0, 1)`.
14    #[default]
15    Sigmoid,
16    /// `tanh(x)`, in `(-1, 1)`.
17    Tanh,
18    /// `max(x, 0)`.
19    ReLU,
20    /// `1` for `x >= 0`, else `0`.
21    BinaryStep,
22    /// `x`, unchanged.
23    Identity,
24}
25
26impl ActivationFunction {
27    /// Every activation function, in declaration order.
28    pub const ALL: [ActivationFunction; 5] = [
29        ActivationFunction::Sigmoid,
30        ActivationFunction::Tanh,
31        ActivationFunction::ReLU,
32        ActivationFunction::BinaryStep,
33        ActivationFunction::Identity,
34    ];
35
36    /// Applies the function to `x`.
37    ///
38    /// ```
39    /// # use only_brain::ActivationFunction;
40    /// assert_eq!(ActivationFunction::ReLU.apply(-2.0), 0.0);
41    /// assert_eq!(ActivationFunction::Identity.apply(-2.0), -2.0);
42    /// ```
43    pub fn apply(self, x: f64) -> f64 {
44        get_activation_function(self)(x)
45    }
46
47    /// The byte that stands for this function in the model file format.
48    ///
49    /// These codes are part of the format and never change; new functions get new
50    /// codes. The first four match the order of the format before 0.3.
51    pub(crate) fn code(self) -> u8 {
52        match self {
53            ActivationFunction::Sigmoid => 0,
54            ActivationFunction::Tanh => 1,
55            ActivationFunction::ReLU => 2,
56            ActivationFunction::BinaryStep => 3,
57            ActivationFunction::Identity => 4,
58        }
59    }
60
61    /// The inverse of [`ActivationFunction::code`].
62    pub(crate) fn from_code(code: u8) -> Option<Self> {
63        Self::ALL.into_iter().find(|function| function.code() == code)
64    }
65}
66
67pub fn sigmoid(x: f64) -> f64 {
68    1. / (1. + (-x).exp())
69}
70
71pub fn tanh(x: f64) -> f64 {
72    x.tanh()
73}
74
75pub fn relu(x: f64) -> f64 {
76    x.max(0.0)
77}
78
79pub fn binary_step(x: f64) -> f64 {
80    if x >= 0.0 { 1.0 } else { 0.0 }
81}
82
83pub fn identity(x: f64) -> f64 {
84    x
85}
86
87/// Returns the plain function behind an [`ActivationFunction`].
88pub fn get_activation_function(func: ActivationFunction) -> fn(f64) -> f64 {
89    match func {
90        ActivationFunction::Sigmoid => sigmoid,
91        ActivationFunction::Tanh => tanh,
92        ActivationFunction::ReLU => relu,
93        ActivationFunction::BinaryStep => binary_step,
94        ActivationFunction::Identity => identity,
95    }
96}
97
98#[cfg(test)]
99mod tests {
100    use super::*;
101
102    const EPSILON: f64 = 1e-12;
103
104    fn assert_close(actual: f64, expected: f64) {
105        assert!(
106            (actual - expected).abs() < EPSILON,
107            "expected {expected}, got {actual}"
108        );
109    }
110
111    #[test]
112    fn sigmoid_is_one_half_at_zero_and_saturates_at_the_extremes() {
113        assert_close(sigmoid(0.0), 0.5);
114        assert_close(sigmoid(f64::INFINITY), 1.0);
115        assert_close(sigmoid(f64::NEG_INFINITY), 0.0);
116    }
117
118    #[test]
119    fn relu_clamps_negatives_to_zero_and_passes_positives_through() {
120        assert_close(relu(-3.5), 0.0);
121        assert_close(relu(0.0), 0.0);
122        assert_close(relu(2.25), 2.25);
123    }
124
125    #[test]
126    fn binary_step_switches_at_zero_inclusive() {
127        assert_close(binary_step(-0.001), 0.0);
128        assert_close(binary_step(0.0), 1.0);
129        assert_close(binary_step(0.001), 1.0);
130    }
131
132    #[test]
133    fn tanh_is_odd_around_zero() {
134        assert_close(tanh(0.0), 0.0);
135        assert_close(tanh(1.3), -tanh(-1.3));
136    }
137
138    /// Guards the bug where a lookup table covered only 3 of the 4 variants and
139    /// panicked on the missing one.
140    #[test]
141    fn every_variant_resolves_to_its_own_function() {
142        let cases = [
143            (ActivationFunction::Sigmoid, sigmoid as fn(f64) -> f64),
144            (ActivationFunction::Tanh, tanh),
145            (ActivationFunction::ReLU, relu),
146            (ActivationFunction::BinaryStep, binary_step),
147            (ActivationFunction::Identity, identity),
148        ];
149        assert_eq!(cases.len(), ActivationFunction::ALL.len());
150
151        for (variant, expected) in cases {
152            let resolved = get_activation_function(variant);
153            for x in [-2.0, -0.5, 0.0, 0.5, 2.0] {
154                assert_close(resolved(x), expected(x));
155            }
156        }
157    }
158
159    /// The codes are written into model files, so they must never move.
160    #[test]
161    fn file_format_codes_are_stable_and_round_trip() {
162        let expected = [
163            (ActivationFunction::Sigmoid, 0),
164            (ActivationFunction::Tanh, 1),
165            (ActivationFunction::ReLU, 2),
166            (ActivationFunction::BinaryStep, 3),
167            (ActivationFunction::Identity, 4),
168        ];
169
170        for (function, code) in expected {
171            assert_eq!(function.code(), code);
172            assert_eq!(ActivationFunction::from_code(code), Some(function));
173        }
174        assert_eq!(ActivationFunction::from_code(5), None);
175    }
176
177    #[test]
178    fn identity_passes_values_through() {
179        assert_close(identity(-3.25), -3.25);
180        assert_close(ActivationFunction::Identity.apply(7.5), 7.5);
181    }
182
183    #[test]
184    fn default_activation_function_is_sigmoid() {
185        assert_eq!(ActivationFunction::default(), ActivationFunction::Sigmoid);
186    }
187}
188