Skip to main content

neural_network_study/
nn.rs

1use crate::matrix::Matrix;
2use rand::{Rng, SeedableRng, rngs::StdRng};
3use serde::{Deserialize, Deserializer, Serialize, de};
4use std::{error::Error, fmt};
5
6fn sigmoid(x: &mut Matrix) {
7    x.apply(|x| 1.0 / (1.0 + (-x).exp()))
8}
9
10fn sigmoid_derivative(x: &mut Matrix) {
11    x.apply(|x| x * (1.0 - x))
12}
13
14fn tanh(x: &mut Matrix) {
15    x.apply(|x| x.tanh())
16}
17
18fn tanh_derivative(x: &mut Matrix) {
19    x.apply(|x| 1.0 - x.powi(2))
20}
21
22fn linear(_: &mut Matrix) {}
23
24fn linear_derivative(x: &mut Matrix) {
25    x.apply(|_| 1.0)
26}
27
28#[derive(Clone, Debug, Serialize, Deserialize, Default)]
29pub enum ActivationFunction {
30    #[default]
31    Sigmoid,
32    Tanh,
33    Linear,
34}
35
36impl ActivationFunction {
37    fn apply(&self, x: &mut Matrix) {
38        match self {
39            ActivationFunction::Sigmoid => sigmoid(x),
40            ActivationFunction::Tanh => tanh(x),
41            ActivationFunction::Linear => linear(x),
42        }
43    }
44
45    fn derivative(&self, x: &mut Matrix) {
46        match self {
47            ActivationFunction::Sigmoid => sigmoid_derivative(x),
48            ActivationFunction::Tanh => tanh_derivative(x),
49            ActivationFunction::Linear => linear_derivative(x),
50        }
51    }
52}
53
54#[derive(Clone, Debug, PartialEq, Eq)]
55pub enum NeuralNetworkError {
56    InvalidLayerCount { got: usize },
57    InvalidLayerSize { layer_index: usize, size: usize },
58    InputLengthMismatch { expected: usize, got: usize },
59    TargetLengthMismatch { expected: usize, got: usize },
60}
61
62impl fmt::Display for NeuralNetworkError {
63    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
64        match self {
65            NeuralNetworkError::InvalidLayerCount { got } => {
66                write!(
67                    f,
68                    "invalid layer count: expected at least 2 layers (input and output), got {got}"
69                )
70            }
71            NeuralNetworkError::InvalidLayerSize { layer_index, size } => {
72                write!(
73                    f,
74                    "invalid layer size at index {layer_index}: expected a positive size, got {size}"
75                )
76            }
77            NeuralNetworkError::InputLengthMismatch { expected, got } => {
78                write!(f, "input length mismatch: expected {expected}, got {got}")
79            }
80            NeuralNetworkError::TargetLengthMismatch { expected, got } => {
81                write!(f, "target length mismatch: expected {expected}, got {got}")
82            }
83        }
84    }
85}
86
87impl Error for NeuralNetworkError {}
88
89/// A simple feedforward neural network with an arbitrary number of layers.
90#[derive(Clone, Debug, Default, Serialize)]
91pub struct NeuralNetwork {
92    layer_sizes: Vec<usize>,
93    weights: Vec<Matrix>,
94    biases: Vec<Matrix>,
95    learning_rate: f64,
96    activation_function: ActivationFunction,
97}
98
99#[derive(Deserialize)]
100struct NeuralNetworkRepr {
101    layer_sizes: Vec<usize>,
102    weights: Vec<Matrix>,
103    biases: Vec<Matrix>,
104    learning_rate: f64,
105    activation_function: ActivationFunction,
106}
107
108impl<'de> Deserialize<'de> for NeuralNetwork {
109    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
110    where
111        D: Deserializer<'de>,
112    {
113        let repr = NeuralNetworkRepr::deserialize(deserializer)?;
114
115        if repr.layer_sizes.len() < 2 {
116            return Err(de::Error::custom(format!(
117                "invalid layer count: expected at least 2 layers, got {}",
118                repr.layer_sizes.len()
119            )));
120        }
121
122        for (layer_index, &size) in repr.layer_sizes.iter().enumerate() {
123            if size == 0 {
124                return Err(de::Error::custom(format!(
125                    "invalid layer size at index {layer_index}: expected a positive size, got {size}"
126                )));
127            }
128        }
129
130        let expected_parameter_layers = repr.layer_sizes.len() - 1;
131        if repr.weights.len() != expected_parameter_layers {
132            return Err(de::Error::custom(format!(
133                "weight layer count mismatch: expected {}, got {}",
134                expected_parameter_layers,
135                repr.weights.len()
136            )));
137        }
138        if repr.biases.len() != expected_parameter_layers {
139            return Err(de::Error::custom(format!(
140                "bias layer count mismatch: expected {}, got {}",
141                expected_parameter_layers,
142                repr.biases.len()
143            )));
144        }
145
146        for (layer_index, pair) in repr.layer_sizes.windows(2).enumerate() {
147            let fan_in = pair[0];
148            let fan_out = pair[1];
149            let weight = &repr.weights[layer_index];
150            let bias = &repr.biases[layer_index];
151
152            if weight.rows() != fan_out || weight.cols() != fan_in {
153                return Err(de::Error::custom(format!(
154                    "weight shape mismatch at layer {layer_index}: expected {}x{}, got {}x{}",
155                    fan_out,
156                    fan_in,
157                    weight.rows(),
158                    weight.cols()
159                )));
160            }
161
162            if bias.rows() != fan_out || bias.cols() != 1 {
163                return Err(de::Error::custom(format!(
164                    "bias shape mismatch at layer {layer_index}: expected {}x1, got {}x{}",
165                    fan_out,
166                    bias.rows(),
167                    bias.cols()
168                )));
169            }
170        }
171
172        Ok(Self {
173            layer_sizes: repr.layer_sizes,
174            weights: repr.weights,
175            biases: repr.biases,
176            learning_rate: repr.learning_rate,
177            activation_function: repr.activation_function,
178        })
179    }
180}
181
182impl NeuralNetwork {
183    fn input_size(&self) -> usize {
184        self.layer_sizes.first().copied().unwrap_or(0)
185    }
186
187    fn output_size(&self) -> usize {
188        self.layer_sizes.last().copied().unwrap_or(0)
189    }
190
191    fn validate_input_len(&self, actual: usize) -> Result<(), NeuralNetworkError> {
192        if actual == self.input_size() {
193            Ok(())
194        } else {
195            Err(NeuralNetworkError::InputLengthMismatch {
196                expected: self.input_size(),
197                got: actual,
198            })
199        }
200    }
201
202    fn validate_target_len(&self, actual: usize) -> Result<(), NeuralNetworkError> {
203        if actual == self.output_size() {
204            Ok(())
205        } else {
206            Err(NeuralNetworkError::TargetLengthMismatch {
207                expected: self.output_size(),
208                got: actual,
209            })
210        }
211    }
212
213    /// Creates a new neural network from a layer-size specification.
214    ///
215    /// `layer_sizes` must include at least input and output layers,
216    /// e.g. `[2, 4, 1]` or `[3, 2]` for a perceptron.
217    /// Weights use Xavier-style initialization and biases start at zero.
218    pub fn new(
219        layer_sizes: Vec<usize>,
220        rng: Option<&mut StdRng>,
221    ) -> Result<Self, NeuralNetworkError> {
222        if layer_sizes.len() < 2 {
223            return Err(NeuralNetworkError::InvalidLayerCount {
224                got: layer_sizes.len(),
225            });
226        }
227
228        for (layer_index, &size) in layer_sizes.iter().enumerate() {
229            if size == 0 {
230                return Err(NeuralNetworkError::InvalidLayerSize { layer_index, size });
231            }
232        }
233
234        let rng = match rng {
235            Some(rng) => rng,
236            None => &mut StdRng::from_os_rng(),
237        };
238
239        let mut weights = Vec::with_capacity(layer_sizes.len() - 1);
240        let mut biases = Vec::with_capacity(layer_sizes.len() - 1);
241
242        for pair in layer_sizes.windows(2) {
243            let fan_in = pair[0];
244            let fan_out = pair[1];
245            let limit = (6.0 / (fan_in + fan_out) as f64).sqrt();
246
247            weights.push(Matrix::random_range(rng, fan_out, fan_in, -limit, limit));
248            biases.push(Matrix::new(fan_out, 1));
249        }
250
251        Ok(NeuralNetwork {
252            layer_sizes,
253            weights,
254            biases,
255            learning_rate: 0.01,
256            activation_function: ActivationFunction::default(),
257        })
258    }
259
260    /// Returns the layer sizes used by this network.
261    pub fn layer_sizes(&self) -> &[usize] {
262        &self.layer_sizes
263    }
264
265    /// Returns the learning rate of the neural network.
266    pub fn learning_rate(&self) -> f64 {
267        self.learning_rate
268    }
269
270    /// Sets the learning rate for the neural network.
271    pub fn set_learning_rate(&mut self, learning_rate: f64) {
272        self.learning_rate = learning_rate;
273    }
274
275    /// Returns the activation function of the neural network.
276    pub fn activation_function(&self) -> &ActivationFunction {
277        &self.activation_function
278    }
279
280    /// Sets the activation function for the neural network.
281    pub fn set_activation_function(&mut self, activation_function: ActivationFunction) {
282        self.activation_function = activation_function;
283    }
284
285    /// Predicts the output for the given input using the neural network.
286    pub fn predict(&self, input: Vec<f64>) -> Result<Vec<f64>, NeuralNetworkError> {
287        self.validate_input_len(input.len())?;
288
289        let mut activation = Matrix::from_col_vec(input);
290
291        for (weights, biases) in self.weights.iter().zip(self.biases.iter()) {
292            let mut layer_input = weights * &activation;
293            layer_input += biases;
294            self.activation_function.apply(&mut layer_input);
295            activation = layer_input;
296        }
297
298        Ok(activation.col(0))
299    }
300
301    /// Trains the neural network using the given input and target output.
302    pub fn train(&mut self, input: Vec<f64>, target: Vec<f64>) -> Result<(), NeuralNetworkError> {
303        self.validate_input_len(input.len())?;
304        self.validate_target_len(target.len())?;
305
306        let mut activations = Vec::with_capacity(self.layer_sizes.len());
307        let mut activation = Matrix::from_col_vec(input);
308        activations.push(activation.clone());
309
310        for (weights, biases) in self.weights.iter().zip(self.biases.iter()) {
311            let mut layer_input = weights * &activation;
312            layer_input += biases;
313            self.activation_function.apply(&mut layer_input);
314            activation = layer_input;
315            activations.push(activation.clone());
316        }
317
318        let target = Matrix::from_col_vec(target);
319        let mut errors = target;
320        errors -= activations
321            .last()
322            .expect("output layer activation is missing");
323
324        for layer_idx in (0..self.weights.len()).rev() {
325            let mut deltas = activations[layer_idx + 1].clone();
326            self.activation_function.derivative(&mut deltas);
327            deltas.hadamard_product(&errors);
328
329            let mut gradients = deltas.clone();
330            gradients *= self.learning_rate;
331
332            let prev_activation_t = activations[layer_idx].transpose();
333            let weight_deltas = &gradients * &prev_activation_t;
334
335            let weights_transposed = self.weights[layer_idx].transpose();
336            let next_errors = &weights_transposed * &deltas;
337
338            self.weights[layer_idx] += &weight_deltas;
339            self.biases[layer_idx] += &gradients;
340
341            errors = next_errors;
342        }
343
344        Ok(())
345    }
346
347    pub fn mutate(&mut self, rng: &mut StdRng, mutation_rate: f64) {
348        for weight in &mut self.weights {
349            for value in weight.data_mut().iter_mut() {
350                if rng.random::<f64>() < mutation_rate {
351                    *value = rng.random_range(-1.0..1.0);
352                }
353            }
354        }
355
356        for bias in &mut self.biases {
357            for value in bias.data_mut().iter_mut() {
358                if rng.random::<f64>() < mutation_rate {
359                    *value = rng.random_range(-1.0..1.0);
360                }
361            }
362        }
363    }
364}
365
366#[cfg(test)]
367pub mod nn_tests {
368    use rand::{SeedableRng, rngs::StdRng};
369    use serde_json;
370
371    fn sigmoid_scalar(value: f64) -> f64 {
372        1.0 / (1.0 + (-value).exp())
373    }
374
375    fn assert_close(actual: f64, expected: f64) {
376        assert!(
377            (actual - expected).abs() < 1e-12,
378            "expected {actual} to be within 1e-12 of {expected}"
379        );
380    }
381
382    #[test]
383    fn it_creates_a_neural_network() {
384        let m = super::NeuralNetwork::new(vec![3, 5, 2], None).unwrap();
385
386        assert_eq!(m.layer_sizes, vec![3, 5, 2]);
387        assert_eq!(m.input_size(), 3);
388        assert_eq!(m.output_size(), 2);
389
390        assert_eq!(m.weights.len(), 2);
391        assert_eq!(m.weights[0].rows(), 5);
392        assert_eq!(m.weights[0].cols(), 3);
393        assert_eq!(m.weights[1].rows(), 2);
394        assert_eq!(m.weights[1].cols(), 5);
395
396        assert_eq!(m.biases.len(), 2);
397        assert_eq!(m.biases[0].rows(), 5);
398        assert_eq!(m.biases[0].cols(), 1);
399        assert_eq!(m.biases[1].rows(), 2);
400        assert_eq!(m.biases[1].cols(), 1);
401    }
402
403    #[test]
404    fn it_creates_a_deep_neural_network() {
405        let m = super::NeuralNetwork::new(vec![3, 4, 4, 2], None).unwrap();
406
407        assert_eq!(m.weights.len(), 3);
408        assert_eq!(m.weights[0].rows(), 4);
409        assert_eq!(m.weights[0].cols(), 3);
410        assert_eq!(m.weights[1].rows(), 4);
411        assert_eq!(m.weights[1].cols(), 4);
412        assert_eq!(m.weights[2].rows(), 2);
413        assert_eq!(m.weights[2].cols(), 4);
414
415        assert_eq!(m.biases.len(), 3);
416        assert_eq!(m.biases[0].rows(), 4);
417        assert_eq!(m.biases[1].rows(), 4);
418        assert_eq!(m.biases[2].rows(), 2);
419    }
420
421    #[test]
422    fn it_creates_a_no_hidden_layer_network() {
423        let m = super::NeuralNetwork::new(vec![3, 2], None).unwrap();
424
425        assert_eq!(m.weights.len(), 1);
426        assert_eq!(m.biases.len(), 1);
427        assert_eq!(m.weights[0].rows(), 2);
428        assert_eq!(m.weights[0].cols(), 3);
429        assert_eq!(m.biases[0].rows(), 2);
430        assert_eq!(m.biases[0].cols(), 1);
431    }
432
433    #[test]
434    pub fn it_predicts() {
435        let m = super::NeuralNetwork::new(vec![3, 5, 2], None).unwrap();
436        let input = vec![0.5, 0.2, 0.1];
437        let output = m.predict(input).unwrap();
438        assert_eq!(output.len(), 2);
439        assert_ne!(output[0], output[1]);
440    }
441
442    #[test]
443    fn predict_handles_deep_and_no_hidden_architectures() {
444        let deep = super::NeuralNetwork::new(vec![3, 4, 4, 2], None).unwrap();
445        let no_hidden = super::NeuralNetwork::new(vec![3, 2], None).unwrap();
446
447        assert_eq!(deep.predict(vec![0.1, 0.2, 0.3]).unwrap().len(), 2);
448        assert_eq!(no_hidden.predict(vec![0.1, 0.2, 0.3]).unwrap().len(), 2);
449    }
450
451    #[test]
452    fn predict_linear_activation_matches_manual_multilayer_math() {
453        let mut nn = super::NeuralNetwork::new(vec![2, 2, 1], None).unwrap();
454        nn.set_activation_function(super::ActivationFunction::Linear);
455
456        nn.weights[0] = super::Matrix::from_vec(2, 2, vec![1.0, 2.0, 3.0, 4.0]);
457        nn.biases[0] = super::Matrix::from_col_vec(vec![0.5, -0.5]);
458        nn.weights[1] = super::Matrix::from_vec(1, 2, vec![2.0, -1.0]);
459        nn.biases[1] = super::Matrix::from_col_vec(vec![1.0]);
460
461        let output = nn.predict(vec![0.25, 0.75]).unwrap();
462        assert!((output[0] - 2.25).abs() < 1e-12);
463    }
464
465    #[test]
466    fn train_updates_all_layers_in_deep_network() {
467        let mut nn = super::NeuralNetwork::new(vec![2, 3, 2, 1], None).unwrap();
468        nn.set_activation_function(super::ActivationFunction::Linear);
469        nn.set_learning_rate(0.1);
470
471        nn.weights[0] = super::Matrix::from_vec(3, 2, vec![0.1, 0.2, 0.3, 0.4, 0.5, 0.6]);
472        nn.weights[1] = super::Matrix::from_vec(2, 3, vec![0.2, 0.1, 0.4, 0.3, 0.5, 0.7]);
473        nn.weights[2] = super::Matrix::from_vec(1, 2, vec![0.9, 0.8]);
474        nn.biases[0] = super::Matrix::from_col_vec(vec![0.0, 0.0, 0.0]);
475        nn.biases[1] = super::Matrix::from_col_vec(vec![0.0, 0.0]);
476        nn.biases[2] = super::Matrix::from_col_vec(vec![0.0]);
477
478        let weights_before: Vec<Vec<f64>> = nn.weights.iter().map(|w| w.data().to_vec()).collect();
479        let biases_before: Vec<Vec<f64>> = nn.biases.iter().map(|b| b.data().to_vec()).collect();
480
481        nn.train(vec![0.9, 0.1], vec![0.2]).unwrap();
482
483        for (idx, weight) in nn.weights.iter().enumerate() {
484            assert_ne!(
485                weight.data(),
486                &weights_before[idx],
487                "expected weights at layer {idx} to change"
488            );
489        }
490        for (idx, bias) in nn.biases.iter().enumerate() {
491            assert_ne!(
492                bias.data(),
493                &biases_before[idx],
494                "expected biases at layer {idx} to change"
495            );
496        }
497    }
498
499    #[test]
500    fn train_uses_downstream_delta_when_updating_hidden_layers() {
501        let mut nn = super::NeuralNetwork::new(vec![1, 1, 1], None).unwrap();
502        nn.set_learning_rate(0.5);
503
504        nn.weights[0] = super::Matrix::from_vec(1, 1, vec![0.5]);
505        nn.biases[0] = super::Matrix::from_col_vec(vec![0.0]);
506        nn.weights[1] = super::Matrix::from_vec(1, 1, vec![-0.4]);
507        nn.biases[1] = super::Matrix::from_col_vec(vec![0.1]);
508
509        let input = 1.0;
510        let target = 0.8;
511        let hidden = sigmoid_scalar(0.5 * input);
512        let output = sigmoid_scalar(-0.4 * hidden + 0.1);
513        let output_error = target - output;
514        let output_delta = output_error * output * (1.0 - output);
515        let hidden_error = -0.4 * output_delta;
516        let hidden_delta = hidden_error * hidden * (1.0 - hidden);
517
518        let expected_hidden_weight = 0.5 + 0.5 * hidden_delta * input;
519        let expected_hidden_bias = 0.0 + 0.5 * hidden_delta;
520        let expected_output_weight = -0.4 + 0.5 * output_delta * hidden;
521        let expected_output_bias = 0.1 + 0.5 * output_delta;
522
523        nn.train(vec![input], vec![target]).unwrap();
524
525        assert_close(nn.weights[0].get(0, 0), expected_hidden_weight);
526        assert_close(nn.biases[0].get(0, 0), expected_hidden_bias);
527        assert_close(nn.weights[1].get(0, 0), expected_output_weight);
528        assert_close(nn.biases[1].get(0, 0), expected_output_bias);
529    }
530
531    #[test]
532    fn mutate_honors_rate_extremes() {
533        let mut rng = StdRng::seed_from_u64(77);
534        let mut nn = super::NeuralNetwork::new(vec![3, 4, 2], Some(&mut rng)).unwrap();
535
536        let original_weights: Vec<Vec<f64>> =
537            nn.weights.iter().map(|w| w.data().to_vec()).collect();
538        let original_biases: Vec<Vec<f64>> = nn.biases.iter().map(|b| b.data().to_vec()).collect();
539
540        nn.mutate(&mut rng, 0.0);
541        for (idx, weight) in nn.weights.iter().enumerate() {
542            assert_eq!(weight.data(), &original_weights[idx]);
543        }
544        for (idx, bias) in nn.biases.iter().enumerate() {
545            assert_eq!(bias.data(), &original_biases[idx]);
546        }
547
548        nn.mutate(&mut rng, 1.0);
549        let mut any_changed = false;
550
551        for (idx, weight) in nn.weights.iter().enumerate() {
552            if weight.data() != original_weights[idx].as_slice() {
553                any_changed = true;
554            }
555            assert!(
556                weight
557                    .data()
558                    .iter()
559                    .all(|value| *value >= -1.0 && *value < 1.0)
560            );
561        }
562        for (idx, bias) in nn.biases.iter().enumerate() {
563            if bias.data() != original_biases[idx].as_slice() {
564                any_changed = true;
565            }
566            assert!(
567                bias.data()
568                    .iter()
569                    .all(|value| *value >= -1.0 && *value < 1.0)
570            );
571        }
572
573        assert!(any_changed, "expected at least one parameter to change");
574    }
575
576    #[test]
577    fn it_learns_the_or_function() {
578        let mut rng = StdRng::seed_from_u64(42);
579        let mut nn = super::NeuralNetwork::new(vec![2, 4, 1], Some(&mut rng)).unwrap();
580        nn.set_learning_rate(0.5);
581
582        let training_data = [
583            (vec![0.0, 0.0], vec![0.0]),
584            (vec![0.0, 1.0], vec![1.0]),
585            (vec![1.0, 0.0], vec![1.0]),
586            (vec![1.0, 1.0], vec![1.0]),
587        ];
588
589        for _ in 0..10_000 {
590            for (input, target) in &training_data {
591                nn.train(input.clone(), target.clone()).unwrap();
592            }
593        }
594
595        assert!(nn.predict(vec![0.0, 0.0]).unwrap()[0] < 0.2);
596        assert!(nn.predict(vec![0.0, 1.0]).unwrap()[0] > 0.8);
597        assert!(nn.predict(vec![1.0, 0.0]).unwrap()[0] > 0.8);
598        assert!(nn.predict(vec![1.0, 1.0]).unwrap()[0] > 0.8);
599    }
600
601    #[test]
602    fn tanh_derivative_uses_activated_output() {
603        let mut x = crate::Matrix::from_col_vec(vec![0.5, -0.25]);
604        super::tanh_derivative(&mut x);
605
606        assert!((x.get(0, 0) - 0.75).abs() < 1e-12);
607        assert!((x.get(1, 0) - 0.9375).abs() < 1e-12);
608    }
609
610    #[test]
611    fn predict_returns_clear_error_for_wrong_input_size() {
612        let nn = super::NeuralNetwork::new(vec![3, 5, 2], None).unwrap();
613
614        assert_eq!(
615            nn.predict(vec![0.1, 0.2]),
616            Err(super::NeuralNetworkError::InputLengthMismatch {
617                expected: 3,
618                got: 2,
619            })
620        );
621    }
622
623    #[test]
624    fn train_returns_clear_error_for_wrong_target_size() {
625        let mut nn = super::NeuralNetwork::new(vec![3, 5, 2], None).unwrap();
626
627        assert_eq!(
628            nn.train(vec![0.1, 0.2, 0.3], vec![1.0]),
629            Err(super::NeuralNetworkError::TargetLengthMismatch {
630                expected: 2,
631                got: 1,
632            })
633        );
634    }
635
636    #[test]
637    fn new_rejects_invalid_layer_vectors() {
638        assert_eq!(
639            super::NeuralNetwork::new(vec![], None).unwrap_err(),
640            super::NeuralNetworkError::InvalidLayerCount { got: 0 }
641        );
642
643        assert_eq!(
644            super::NeuralNetwork::new(vec![3], None).unwrap_err(),
645            super::NeuralNetworkError::InvalidLayerCount { got: 1 }
646        );
647
648        assert_eq!(
649            super::NeuralNetwork::new(vec![0, 5, 2], None).unwrap_err(),
650            super::NeuralNetworkError::InvalidLayerSize {
651                layer_index: 0,
652                size: 0,
653            }
654        );
655
656        assert_eq!(
657            super::NeuralNetwork::new(vec![3, 0, 2], None).unwrap_err(),
658            super::NeuralNetworkError::InvalidLayerSize {
659                layer_index: 1,
660                size: 0,
661            }
662        );
663
664        assert_eq!(
665            super::NeuralNetwork::new(vec![3, 5, 0], None).unwrap_err(),
666            super::NeuralNetworkError::InvalidLayerSize {
667                layer_index: 2,
668                size: 0,
669            }
670        );
671    }
672
673    #[test]
674    fn new_uses_zero_biases() {
675        let nn = super::NeuralNetwork::new(vec![3, 5, 4, 2], None).unwrap();
676
677        assert!(
678            nn.biases
679                .iter()
680                .all(|bias| bias.data().iter().all(|value| *value == 0.0))
681        );
682    }
683
684    #[test]
685    fn new_uses_xavier_weight_ranges() {
686        let mut rng = StdRng::seed_from_u64(7);
687        let layer_sizes = vec![3, 5, 4, 2];
688        let nn = super::NeuralNetwork::new(layer_sizes.clone(), Some(&mut rng)).unwrap();
689
690        for (weight, pair) in nn.weights.iter().zip(layer_sizes.windows(2)) {
691            let fan_in = pair[0] as f64;
692            let fan_out = pair[1] as f64;
693            let limit = (6.0_f64 / (fan_in + fan_out)).sqrt();
694
695            assert!(
696                weight
697                    .data()
698                    .iter()
699                    .all(|value| *value >= -limit && *value < limit)
700            );
701        }
702    }
703
704    #[test]
705    fn it_learns_the_xor_function() {
706        let mut rng = StdRng::seed_from_u64(99);
707        let mut nn = super::NeuralNetwork::new(vec![2, 4, 1], Some(&mut rng)).unwrap();
708        nn.set_learning_rate(0.5);
709
710        let training_data = [
711            (vec![0.0, 0.0], vec![0.0]),
712            (vec![0.0, 1.0], vec![1.0]),
713            (vec![1.0, 0.0], vec![1.0]),
714            (vec![1.0, 1.0], vec![0.0]),
715        ];
716
717        for _ in 0..20_000 {
718            for (input, target) in &training_data {
719                nn.train(input.clone(), target.clone()).unwrap();
720            }
721        }
722
723        assert!(nn.predict(vec![0.0, 0.0]).unwrap()[0] < 0.2);
724        assert!(nn.predict(vec![0.0, 1.0]).unwrap()[0] > 0.8);
725        assert!(nn.predict(vec![1.0, 0.0]).unwrap()[0] > 0.8);
726        assert!(nn.predict(vec![1.0, 1.0]).unwrap()[0] < 0.2);
727    }
728
729    #[test]
730    fn it_learns_the_xor_function_with_deeper_network() {
731        let mut rng = StdRng::seed_from_u64(314);
732        let mut nn = super::NeuralNetwork::new(vec![2, 4, 4, 1], Some(&mut rng)).unwrap();
733        nn.set_learning_rate(0.5);
734
735        let training_data = [
736            (vec![0.0, 0.0], vec![0.0]),
737            (vec![0.0, 1.0], vec![1.0]),
738            (vec![1.0, 0.0], vec![1.0]),
739            (vec![1.0, 1.0], vec![0.0]),
740        ];
741
742        for _ in 0..20_000 {
743            for (input, target) in &training_data {
744                nn.train(input.clone(), target.clone()).unwrap();
745            }
746        }
747
748        assert!(nn.predict(vec![0.0, 0.0]).unwrap()[0] < 0.2);
749        assert!(nn.predict(vec![0.0, 1.0]).unwrap()[0] > 0.8);
750        assert!(nn.predict(vec![1.0, 0.0]).unwrap()[0] > 0.8);
751        assert!(nn.predict(vec![1.0, 1.0]).unwrap()[0] < 0.2);
752    }
753
754    #[test]
755    fn perceptron_architecture_learns_linearly_separable_data() {
756        let mut rng = StdRng::seed_from_u64(202);
757        let mut nn = super::NeuralNetwork::new(vec![2, 1], Some(&mut rng)).unwrap();
758        nn.set_learning_rate(0.5);
759
760        let training_data = [
761            (vec![0.0, 0.0], vec![0.0]),
762            (vec![0.0, 1.0], vec![0.0]),
763            (vec![1.0, 0.0], vec![1.0]),
764            (vec![1.0, 1.0], vec![1.0]),
765        ];
766
767        for _ in 0..12_000 {
768            for (input, target) in &training_data {
769                nn.train(input.clone(), target.clone()).unwrap();
770            }
771        }
772
773        assert!(nn.predict(vec![0.0, 0.0]).unwrap()[0] < 0.2);
774        assert!(nn.predict(vec![0.0, 1.0]).unwrap()[0] < 0.2);
775        assert!(nn.predict(vec![1.0, 0.0]).unwrap()[0] > 0.8);
776        assert!(nn.predict(vec![1.0, 1.0]).unwrap()[0] > 0.8);
777    }
778
779    #[test]
780    fn serde_round_trip_preserves_predictions() {
781        let mut rng = StdRng::seed_from_u64(123);
782        let mut nn = super::NeuralNetwork::new(vec![2, 4, 1], Some(&mut rng)).unwrap();
783        nn.set_learning_rate(0.5);
784
785        let training_data = [
786            (vec![0.0, 0.0], vec![0.0]),
787            (vec![0.0, 1.0], vec![1.0]),
788            (vec![1.0, 0.0], vec![1.0]),
789            (vec![1.0, 1.0], vec![0.0]),
790        ];
791
792        for _ in 0..5_000 {
793            for (input, target) in &training_data {
794                nn.train(input.clone(), target.clone()).unwrap();
795            }
796        }
797
798        let probe_input = vec![0.25, 0.75];
799        let before = nn.predict(probe_input.clone()).unwrap();
800
801        let json = serde_json::to_string(&nn).unwrap();
802        let restored: super::NeuralNetwork = serde_json::from_str(&json).unwrap();
803        let after = restored.predict(probe_input).unwrap();
804
805        assert_eq!(before, after);
806    }
807
808    #[test]
809    fn serde_rejects_networks_with_invalid_matrix_shapes() {
810        let json = r#"{
811            "layer_sizes": [2, 1],
812            "weights": [
813                { "rows": 1, "cols": 2, "data": [0.25] }
814            ],
815            "biases": [
816                { "rows": 1, "cols": 1, "data": [0.0] }
817            ],
818            "learning_rate": 0.1,
819            "activation_function": "Sigmoid"
820        }"#;
821
822        let result = serde_json::from_str::<super::NeuralNetwork>(json);
823
824        assert!(
825            result.is_err(),
826            "deserialization should reject matrix data whose length does not match rows * cols"
827        );
828    }
829
830    #[test]
831    fn serde_rejects_networks_whose_shapes_do_not_match_layer_sizes() {
832        let json = r#"{
833            "layer_sizes": [2, 1],
834            "weights": [
835                { "rows": 1, "cols": 1, "data": [0.25] }
836            ],
837            "biases": [
838                { "rows": 1, "cols": 1, "data": [0.0] }
839            ],
840            "learning_rate": 0.1,
841            "activation_function": "Sigmoid"
842        }"#;
843
844        let result = serde_json::from_str::<super::NeuralNetwork>(json);
845
846        assert!(
847            result.is_err(),
848            "deserialization should reject weights whose shape does not match layer_sizes"
849        );
850    }
851}