only_brain/
activation_functions.rs1use serde::{Deserialize, Serialize};
2
3#[derive(Clone, Copy, PartialEq, Eq, Hash, Debug, Default, Serialize, Deserialize)]
11#[non_exhaustive]
12pub enum ActivationFunction {
13 #[default]
15 Sigmoid,
16 Tanh,
18 ReLU,
20 BinaryStep,
22 Identity,
24}
25
26impl ActivationFunction {
27 pub const ALL: [ActivationFunction; 5] = [
29 ActivationFunction::Sigmoid,
30 ActivationFunction::Tanh,
31 ActivationFunction::ReLU,
32 ActivationFunction::BinaryStep,
33 ActivationFunction::Identity,
34 ];
35
36 pub fn apply(self, x: f64) -> f64 {
44 get_activation_function(self)(x)
45 }
46
47 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 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
87pub 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 #[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 #[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