Skip to main content

only_brain/
perceptron.rs

1use crate::{
2    activation_functions::{get_activation_function, ActivationFunction},
3    BVector,
4};
5use std::fmt;
6
7/// Perceptron
8///
9/// This is a single perceptron implementation. It contains weights, bias and an activation function.
10/// You can use this struct and its methods to create, manipulate and even implement your ways to
11/// train a perceptron (find the weights).
12/// # Example
13/// ```
14/// use only_brain::Perceptron;
15/// use only_brain::ActivationFunction;
16/// use only_brain::bvector;
17///
18/// fn main() {
19///     let mut perceptron = Perceptron::<2>::new(ActivationFunction::Sigmoid);
20///
21///     perceptron.set_weights(bvector![0.5, -0.2]);
22///     perceptron.set_bias(0.3);
23///     println!("{}", perceptron);
24///
25///     let inputs = bvector![0.6, 0.4];
26///     let output = perceptron.feed_forward(&inputs);
27///     println!("Output: {}", output);
28/// }
29/// ```
30#[derive(Debug, Clone, PartialEq)]
31pub struct Perceptron<const N: usize> {
32    weights: BVector<f64, N>,
33    bias: f64,
34    activation_function: ActivationFunction,
35}
36
37impl<const N: usize> Perceptron<N> {
38    pub fn new(activation_function: ActivationFunction) -> Self {
39        let weights = BVector::<f64, N>::from_element(0.0);
40        let bias = 0.0;
41
42        Self {
43            weights,
44            bias,
45            activation_function,
46        }
47    }
48
49    pub fn set_weights(&mut self, weights: BVector<f64, N>) {
50        self.weights = weights;
51    }
52
53    pub fn set_bias(&mut self, bias: f64) {
54        self.bias = bias;
55    }
56
57    pub fn feed_forward(&self, inputs: &BVector<f64, N>) -> f64 {
58        let weighted_sum = self.weights.dot(inputs) + self.bias;
59        get_activation_function(self.activation_function)(weighted_sum)
60    }
61}
62
63impl<const N: usize> fmt::Display for Perceptron<N> {
64    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
65        writeln!(f, "Perceptron:")?;
66        writeln!(f, "Activation Function: {:?}", self.activation_function)?;
67        writeln!(f)?;
68        writeln!(f, "Inputs Size: {}", N)?;
69        writeln!(f, "Weights: {:?}", self.weights)?;
70        writeln!(f, "Bias: {}", self.bias)?;
71        writeln!(f)?;
72
73        Ok(())
74    }
75}
76
77#[cfg(test)]
78mod tests {
79    use super::*;
80    use crate::bvector;
81
82    const EPSILON: f64 = 1e-12;
83
84    #[test]
85    fn a_new_perceptron_has_zero_weights_and_bias() {
86        let perceptron = Perceptron::<3>::new(ActivationFunction::Identity);
87
88        assert_eq!(perceptron.feed_forward(&bvector![7.0, -2.0, 3.0]), 0.0);
89    }
90
91    #[test]
92    fn feed_forward_applies_weights_bias_and_activation() {
93        let mut perceptron = Perceptron::<2>::new(ActivationFunction::Identity);
94        perceptron.set_weights(bvector![0.5, -0.25]);
95        perceptron.set_bias(0.1);
96
97        // 0.5 * 2 - 0.25 * 4 + 0.1
98        assert!((perceptron.feed_forward(&bvector![2.0, 4.0]) - 0.1).abs() < EPSILON);
99
100        let mut perceptron = Perceptron::<2>::new(ActivationFunction::Sigmoid);
101        perceptron.set_weights(bvector![0.5, -0.25]);
102        perceptron.set_bias(0.1);
103        let expected = 1.0 / (1.0 + (-0.1f64).exp());
104        assert!((perceptron.feed_forward(&bvector![2.0, 4.0]) - expected).abs() < EPSILON);
105    }
106
107    #[test]
108    fn a_binary_step_perceptron_computes_or() {
109        let mut or_gate = Perceptron::<2>::new(ActivationFunction::BinaryStep);
110        or_gate.set_weights(bvector![1.0, 1.0]);
111        or_gate.set_bias(-0.5);
112
113        for (inputs, expected) in [([0.0, 0.0], 0.0), ([0.0, 1.0], 1.0), ([1.0, 0.0], 1.0), ([1.0, 1.0], 1.0)] {
114            assert_eq!(or_gate.feed_forward(&BVector::from_array(inputs)), expected, "OR{inputs:?}");
115        }
116    }
117
118    #[test]
119    fn display_shows_activation_size_and_bias() {
120        let mut perceptron = Perceptron::<4>::new(ActivationFunction::Tanh);
121        perceptron.set_bias(0.75);
122
123        let text = perceptron.to_string();
124
125        assert!(text.contains("Activation Function: Tanh"), "{text}");
126        assert!(text.contains("Inputs Size: 4"), "{text}");
127        assert!(text.contains("Bias: 0.75"), "{text}");
128    }
129}