1use crate::{
2 activation_functions::{get_activation_function, ActivationFunction},
3 BVector,
4};
5use std::fmt;
6
7#[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 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}