1use crate::activation_functions::ActivationFunction;
2use crate::layer::Layer;
3use crate::ModelError;
4use rand::{rng, Rng};
5use std::fmt;
6use serde::{Deserialize, Serialize};
7
8#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
83#[serde(into = "NetworkData", try_from = "NetworkData")]
84pub struct NeuralNetwork<const IN: usize, const OUT: usize> {
85 layers: Vec<Layer>,
86}
87
88#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
94pub(crate) struct NetworkData {
95 pub(crate) layers: Vec<LayerData>,
96}
97
98#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
100pub(crate) struct LayerData {
101 pub(crate) activation: ActivationFunction,
102 pub(crate) weights: Vec<Vec<f64>>,
104 pub(crate) biases: Vec<f64>,
105}
106
107impl<const IN: usize, const OUT: usize> NeuralNetwork<IN, OUT> {
108 pub fn new(hidden: &[usize]) -> Self {
132 Self::new_with_rng(hidden, &mut rng())
133 }
134
135 pub fn new_with_rng<R: Rng>(hidden: &[usize], rng: &mut R) -> Self {
144 Self::build(hidden, |neurons, inputs| {
145 Layer::random(neurons, inputs, ActivationFunction::default(), rng)
146 })
147 }
148
149 pub fn from_parameters(hidden: &[usize], parameters: &[f64]) -> Self {
174 let mut network = Self::build(hidden, |neurons, inputs| {
175 Layer::zeros(neurons, inputs, ActivationFunction::default())
176 });
177 network.set_parameters(parameters);
178 network
179 }
180
181 fn build(hidden: &[usize], mut make: impl FnMut(usize, usize) -> Layer) -> Self {
183 let layers = Self::layer_shapes(hidden)
184 .map(|(neurons, inputs)| make(neurons, inputs))
185 .collect();
186
187 Self { layers }
188 }
189
190 fn layer_shapes(hidden: &[usize]) -> impl Iterator<Item = (usize, usize)> + '_ {
192 const { assert!(IN > 0, "a network needs at least one input neuron") };
193 const { assert!(OUT > 0, "a network needs at least one output neuron") };
194 assert!(
195 hidden.iter().all(|&size| size > 0),
196 "every hidden layer must have at least one neuron, got {hidden:?}"
197 );
198
199 let inputs = std::iter::once(IN).chain(hidden.iter().copied());
200 let neurons = hidden.iter().copied().chain(std::iter::once(OUT));
201 neurons.zip(inputs)
202 }
203
204 pub fn feed_forward(&self, inputs: &[f64; IN]) -> [f64; OUT] {
221 let widest = self.layers.iter().map(Layer::size).fold(IN, usize::max);
224 let mut scratch = vec![0.0; widest * 2];
225 let (mut current, mut next) = scratch.split_at_mut(widest);
226
227 current[..IN].copy_from_slice(inputs);
228 let mut width = IN;
229
230 for layer in &self.layers {
231 layer.forward_into(¤t[..width], &mut next[..layer.size()]);
232 width = layer.size();
233 std::mem::swap(&mut current, &mut next);
234 }
235
236 <[f64; OUT]>::try_from(¤t[..width])
239 .expect("output layer width should match the OUT type parameter")
240 }
241
242 pub fn parameter_count(&self) -> usize {
245 self.layers.iter().map(Layer::parameter_count).sum()
246 }
247
248 pub fn parameter_count_for(hidden: &[usize]) -> usize {
263 Self::layer_shapes(hidden)
264 .map(|(neurons, inputs)| Layer::parameter_count_for(neurons, inputs))
265 .sum()
266 }
267
268 pub fn parameters(&self) -> Vec<f64> {
271 let mut parameters = Vec::with_capacity(self.parameter_count());
272 for layer in &self.layers {
273 layer.extend_parameters(&mut parameters);
274 }
275 parameters
276 }
277
278 pub fn set_parameters(&mut self, parameters: &[f64]) {
286 let expected = self.parameter_count();
287 assert_eq!(
288 parameters.len(),
289 expected,
290 "Incompatible parameter count: expected {expected} parameters, got {}",
291 parameters.len()
292 );
293
294 let mut rest = parameters;
295 for layer in &mut self.layers {
296 rest = layer.take_parameters(rest);
297 }
298 }
299
300 pub fn set_layer_weights<R: AsRef<[f64]>>(&mut self, layer: usize, weights: &[R]) {
327 self.layer_mut(layer).set_weights(weights);
328 }
329
330 pub fn layer_weights(&self, layer: usize) -> Vec<Vec<f64>> {
337 self.layer(layer).weight_rows()
338 }
339
340 pub fn set_layer_biases(&mut self, layer: usize, biases: &[f64]) {
347 self.layer_mut(layer).set_biases(biases);
348 }
349
350 pub fn layer_biases(&self, layer: usize) -> &[f64] {
356 self.layer(layer).biases().as_slice()
357 }
358
359 pub fn set_weight(&mut self, layer: usize, neuron: usize, input: usize, weight: f64) {
362 self.layer_mut(layer).set_weight(neuron, input, weight);
363 }
364
365 pub fn get_weight(&self, layer: usize, neuron: usize, input: usize) -> f64 {
368 self.layer(layer).weights()[(neuron, input)]
369 }
370
371 pub fn set_bias(&mut self, layer: usize, neuron: usize, bias: f64) {
374 self.layer_mut(layer).set_bias(neuron, bias);
375 }
376
377 pub fn get_bias(&self, layer: usize, neuron: usize) -> f64 {
380 self.layer(layer).biases()[neuron]
381 }
382
383 fn layer(&self, layer: usize) -> &Layer {
385 if layer == 0 {
386 panic!("Invalid layer index");
387 }
388 &self.layers[layer - 1]
389 }
390
391 pub(crate) fn layers(&self) -> &[Layer] {
393 &self.layers
394 }
395
396 fn layer_mut(&mut self, layer: usize) -> &mut Layer {
397 if layer == 0 {
398 panic!("Invalid layer index");
399 }
400 &mut self.layers[layer - 1]
401 }
402
403 pub fn num_layers(&self) -> usize {
405 self.layers.len() + 1
406 }
407
408 pub fn layer_size(&self, layer: usize) -> usize {
410 if layer == 0 {
411 return IN;
412 }
413 self.layer(layer).size()
414 }
415
416 pub fn hidden_layer_sizes(&self) -> Vec<usize> {
427 let hidden = &self.layers[..self.layers.len() - 1];
428 hidden.iter().map(Layer::size).collect()
429 }
430
431 pub fn layer_activation(&self, layer: usize) -> ActivationFunction {
438 self.layer(layer).activation()
439 }
440
441 pub fn set_layer_activation(&mut self, layer: usize, activation_function: ActivationFunction) {
448 self.layer_mut(layer).set_activation(activation_function);
449 }
450
451 pub fn set_activation_function(&mut self, activation_function: ActivationFunction) {
468 for layer in &mut self.layers {
469 layer.set_activation(activation_function);
470 }
471 }
472
473 pub fn set_output_activation(&mut self, activation_function: ActivationFunction) {
475 self.set_layer_activation(self.num_layers() - 1, activation_function);
476 }
477
478 pub fn output_activation(&self) -> ActivationFunction {
480 self.layer_activation(self.num_layers() - 1)
481 }
482
483 pub fn print(&self) {
484 for layer in &self.layers {
485 println!("{:?} {} {}", layer.activation(), layer.weights(), layer.biases());
486 }
487 }
488}
489
490impl<const IN: usize, const OUT: usize> From<NeuralNetwork<IN, OUT>> for NetworkData {
491 fn from(network: NeuralNetwork<IN, OUT>) -> Self {
492 NetworkData::from(&network)
493 }
494}
495
496impl<const IN: usize, const OUT: usize> From<&NeuralNetwork<IN, OUT>> for NetworkData {
497 fn from(network: &NeuralNetwork<IN, OUT>) -> Self {
498 let layers = network
499 .layers
500 .iter()
501 .map(|layer| LayerData {
502 activation: layer.activation(),
503 weights: layer.weight_rows(),
504 biases: layer.biases().iter().copied().collect(),
505 })
506 .collect();
507
508 NetworkData { layers }
509 }
510}
511
512impl<const IN: usize, const OUT: usize> TryFrom<NetworkData> for NeuralNetwork<IN, OUT> {
513 type Error = ModelError;
514
515 fn try_from(data: NetworkData) -> Result<Self, Self::Error> {
518 let first = data.layers.first().ok_or(ModelError::EmptyNetwork)?;
519 if let Some(row) = first.weights.first() {
520 if row.len() != IN {
521 return Err(ModelError::DimensionMismatch {
522 end: "input",
523 expected: IN,
524 found: row.len(),
525 });
526 }
527 }
528
529 let mut inputs = IN;
530 let mut layers = Vec::with_capacity(data.layers.len());
531 for (index, layer) in data.layers.iter().enumerate() {
532 let layer = Layer::from_parts(&layer.weights, &layer.biases, layer.activation, inputs)
533 .map_err(|reason| ModelError::InconsistentLayer {
534 layer: index + 1,
535 reason,
536 })?;
537 inputs = layer.size();
538 layers.push(layer);
539 }
540
541 if inputs != OUT {
542 return Err(ModelError::DimensionMismatch {
543 end: "output",
544 expected: OUT,
545 found: inputs,
546 });
547 }
548
549 Ok(Self { layers })
550 }
551}
552
553impl<const IN: usize, const OUT: usize> fmt::Display for NeuralNetwork<IN, OUT> {
554 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
555 writeln!(f, "Neural Network")?;
556 writeln!(f)?;
557 writeln!(f, "Input Layer Size: {IN}")?;
558 writeln!(f)?;
559 for (index, layer) in self.layers.iter().enumerate() {
560 writeln!(f, "Layer {}: {}", index + 1, layer)?;
561 }
562 Ok(())
563 }
564}
565
566#[cfg(test)]
567mod tests {
568 use super::*;
569 use crate::activation_functions::{binary_step, identity, relu, sigmoid, tanh};
570
571 const EPSILON: f64 = 1e-12;
572
573 fn assert_all_close(actual: &[f64], expected: &[f64]) {
574 assert_eq!(actual.len(), expected.len(), "length mismatch");
575 for (i, (a, e)) in actual.iter().zip(expected).enumerate() {
576 assert!((a - e).abs() < EPSILON, "at index {i}: expected {e}, got {a}");
577 }
578 }
579
580 fn fixed_network() -> NeuralNetwork<2, 1> {
582 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
583 nn.set_layer_weights(1, &[[0.5, -0.25]]);
584 nn.set_layer_biases(1, &[0.1]);
585 nn
586 }
587
588 #[test]
589 fn new_reports_layer_count_and_sizes() {
590 let nn = NeuralNetwork::<3, 2>::new(&[5]);
591
592 assert_eq!(nn.num_layers(), 3);
593 assert_eq!(nn.layer_size(0), 3);
594 assert_eq!(nn.layer_size(1), 5);
595 assert_eq!(nn.layer_size(2), 2);
596 }
597
598 #[test]
601 fn new_without_hidden_layers_wires_input_straight_to_output() {
602 let nn = NeuralNetwork::<3, 2>::new(&[]);
603
604 assert_eq!(nn.num_layers(), 2);
605 assert_eq!(nn.layer_size(0), 3);
606 assert_eq!(nn.layer_size(1), 2);
607 }
608
609 #[test]
610 #[should_panic(expected = "at least one neuron")]
611 fn new_rejects_a_zero_sized_hidden_layer() {
612 NeuralNetwork::<2, 1>::new(&[0]);
613 }
614
615 #[test]
616 fn new_with_rng_is_reproducible_for_the_same_seed() {
617 use rand::SeedableRng;
618
619 let mut first_rng = rand::rngs::StdRng::seed_from_u64(42);
620 let mut second_rng = rand::rngs::StdRng::seed_from_u64(42);
621
622 let first = NeuralNetwork::<3, 2>::new_with_rng(&[4], &mut first_rng);
623 let second = NeuralNetwork::<3, 2>::new_with_rng(&[4], &mut second_rng);
624
625 for layer in 1..first.num_layers() {
626 for neuron in 0..first.layer_size(layer) {
627 for input in 0..first.layer_size(layer - 1) {
628 assert_eq!(
629 first.get_weight(layer, neuron, input),
630 second.get_weight(layer, neuron, input),
631 "weight differs at layer {layer}, neuron {neuron}, input {input}"
632 );
633 }
634 }
635 }
636 }
637
638 #[test]
639 fn feed_forward_applies_weights_bias_and_activation() {
640 let nn = fixed_network();
641
642 let output = nn.feed_forward(&[1.0, 2.0]);
644
645 assert_all_close(&output, &[sigmoid(0.1)]);
646 }
647
648 #[test]
649 fn feed_forward_defaults_to_sigmoid() {
650 let nn = fixed_network();
651
652 assert_eq!(nn.layer_activation(1), ActivationFunction::Sigmoid);
653 assert_all_close(&nn.feed_forward(&[1.0, 2.0]), &[sigmoid(0.1)]);
654 }
655
656 #[test]
659 fn feed_forward_honours_every_activation_function() {
660 let cases = [
661 (ActivationFunction::Sigmoid, sigmoid as fn(f64) -> f64),
662 (ActivationFunction::Tanh, tanh),
663 (ActivationFunction::ReLU, relu),
664 (ActivationFunction::BinaryStep, binary_step),
665 (ActivationFunction::Identity, identity),
666 ];
667
668 for (variant, expected) in cases {
669 let mut nn = fixed_network();
670 nn.set_activation_function(variant);
671
672 assert_eq!(nn.layer_activation(1), variant);
673 assert_all_close(&nn.feed_forward(&[1.0, 2.0]), &[expected(0.1)]);
674 }
675 }
676
677 #[test]
679 fn binary_step_does_not_panic() {
680 let mut nn = fixed_network();
681 nn.set_activation_function(ActivationFunction::BinaryStep);
682
683 assert_all_close(&nn.feed_forward(&[1.0, 2.0]), &[1.0]);
684 assert_all_close(&nn.feed_forward(&[-1.0, 2.0]), &[0.0]);
685 }
686
687 #[test]
688 fn set_and_get_weight_round_trip() {
689 let mut nn = NeuralNetwork::<2, 2>::new(&[]);
690 nn.set_weight(1, 1, 0, 0.75);
691
692 assert_eq!(nn.get_weight(1, 1, 0), 0.75);
693 }
694
695 #[test]
696 #[should_panic(expected = "Invalid layer index")]
697 fn set_layer_weights_rejects_layer_zero() {
698 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
699 nn.set_layer_weights(0, &[[0.5, 0.5]]);
700 }
701
702 #[test]
703 #[should_panic(expected = "Incompatible weights matrix size")]
704 fn set_layer_weights_rejects_a_mismatched_matrix() {
705 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
706 nn.set_layer_weights(1, &[[0.5, 0.5, 0.5]]);
707 }
708
709 #[test]
710 #[should_panic(expected = "Incompatible biases vector size")]
711 fn set_layer_biases_rejects_a_mismatched_vector() {
712 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
713 nn.set_layer_biases(1, &[0.1, 0.2]);
714 }
715
716 #[test]
718 fn display_names_the_activation_function() {
719 let mut nn = fixed_network();
720 nn.set_activation_function(ActivationFunction::ReLU);
721
722 let rendered = nn.to_string();
723
724 assert!(
725 rendered.contains("Activation Function: ReLU"),
726 "unexpected output:\n{rendered}"
727 );
728 assert!(
729 !rendered.contains("0x"),
730 "output leaked a pointer address:\n{rendered}"
731 );
732 }
733
734 #[test]
735 fn display_reports_the_input_layer_size() {
736 let nn = NeuralNetwork::<4, 1>::new(&[2]);
737
738 assert!(nn.to_string().contains("Input Layer Size: 4"));
739 }
740
741 #[test]
742 fn layer_weights_and_biases_round_trip_in_row_major_order() {
743 let mut nn = NeuralNetwork::<2, 3>::new(&[]);
744 let weights = [[0.1, 0.2], [0.3, 0.4], [0.5, 0.6]];
745
746 nn.set_layer_weights(1, &weights);
747 nn.set_layer_biases(1, &[0.7, 0.8, 0.9]);
748
749 assert_eq!(nn.layer_weights(1), weights.map(Vec::from).to_vec());
750 assert_eq!(nn.layer_biases(1), &[0.7, 0.8, 0.9]);
751 assert_eq!(nn.get_weight(1, 2, 0), 0.5);
753 }
754
755 #[test]
756 fn set_layer_weights_accepts_weights_built_at_runtime() {
757 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
758 let weights: Vec<Vec<f64>> = vec![vec![0.5, -0.25]];
759
760 nn.set_layer_weights(1, &weights);
761
762 assert_eq!(nn.layer_weights(1), weights);
763 }
764
765 #[test]
766 #[should_panic(expected = "Incompatible weights matrix size")]
767 fn set_layer_weights_rejects_ragged_rows() {
768 let mut nn = NeuralNetwork::<2, 2>::new(&[]);
769 let ragged: Vec<Vec<f64>> = vec![vec![0.1, 0.2], vec![0.3]];
770
771 nn.set_layer_weights(1, &ragged);
772 }
773
774 #[test]
775 #[should_panic(expected = "Incompatible weights matrix size")]
776 fn set_layer_weights_rejects_the_wrong_number_of_rows() {
777 let mut nn = NeuralNetwork::<2, 2>::new(&[]);
778 nn.set_layer_weights(1, &[[0.1, 0.2]]);
779 }
780
781 #[test]
782 fn set_and_get_bias_round_trip() {
783 let mut nn = NeuralNetwork::<2, 2>::new(&[]);
784 nn.set_bias(1, 1, -0.3);
785
786 assert_eq!(nn.get_bias(1, 1), -0.3);
787 assert_eq!(nn.layer_biases(1), &[0.0, -0.3]);
788 }
789
790 #[test]
791 #[should_panic(expected = "Invalid layer index")]
792 fn layer_biases_rejects_layer_zero() {
793 let nn = NeuralNetwork::<2, 1>::new(&[]);
794 nn.layer_biases(0);
795 }
796
797 fn numbered_network() -> NeuralNetwork<2, 1> {
799 let mut nn = NeuralNetwork::<2, 1>::new(&[2]);
800 nn.set_layer_weights(1, &[[1.0, 2.0],
801 [3.0, 4.0]]);
802 nn.set_layer_biases(1, &[5.0, 6.0]);
803 nn.set_layer_weights(2, &[[7.0, 8.0]]);
804 nn.set_layer_biases(2, &[9.0]);
805 nn
806 }
807
808 #[test]
809 fn hidden_and_output_layers_can_use_different_activations() {
810 let mut nn = NeuralNetwork::<1, 1>::new(&[1]);
811 nn.set_layer_weights(1, &[[1.0]]);
812 nn.set_layer_weights(2, &[[1.0]]);
813 nn.set_activation_function(ActivationFunction::ReLU);
814 nn.set_output_activation(ActivationFunction::Tanh);
815
816 assert_all_close(&nn.feed_forward(&[-2.0]), &[0.0]);
818 assert_all_close(&nn.feed_forward(&[2.0]), &[tanh(2.0)]);
819 assert_eq!(nn.layer_activation(1), ActivationFunction::ReLU);
820 assert_eq!(nn.output_activation(), ActivationFunction::Tanh);
821 }
822
823 #[test]
824 fn set_layer_activation_changes_only_that_layer() {
825 let mut nn = NeuralNetwork::<2, 1>::new(&[3, 3]);
826 nn.set_layer_activation(2, ActivationFunction::Identity);
827
828 assert_eq!(nn.layer_activation(1), ActivationFunction::Sigmoid);
829 assert_eq!(nn.layer_activation(2), ActivationFunction::Identity);
830 assert_eq!(nn.layer_activation(3), ActivationFunction::Sigmoid);
831 }
832
833 #[test]
834 #[should_panic(expected = "Invalid layer index")]
835 fn the_input_layer_has_no_activation() {
836 NeuralNetwork::<2, 1>::new(&[]).layer_activation(0);
837 }
838
839 #[test]
840 fn parameters_follow_the_documented_order() {
841 let nn = numbered_network();
842
843 assert_eq!(
844 nn.parameters(),
845 vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0]
846 );
847 }
848
849 #[test]
850 fn parameter_count_matches_the_shape() {
851 let nn = NeuralNetwork::<10, 3>::new(&[8]);
852
853 assert_eq!(nn.parameter_count(), 11 * 8 + 9 * 3);
854 assert_eq!(nn.parameters().len(), nn.parameter_count());
855 assert_eq!(NeuralNetwork::<10, 3>::parameter_count_for(&[8]), nn.parameter_count());
856 assert_eq!(NeuralNetwork::<10, 3>::parameter_count_for(&[]), 11 * 3);
857 }
858
859 #[test]
860 fn set_parameters_writes_the_documented_order() {
861 let mut nn = NeuralNetwork::<2, 1>::new(&[2]);
862 nn.set_parameters(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0]);
863
864 assert_eq!(nn, numbered_network());
865 assert_eq!(nn.get_weight(1, 1, 0), 3.0);
866 assert_eq!(nn.get_bias(2, 0), 9.0);
867 }
868
869 #[test]
870 fn set_parameters_keeps_activation_functions() {
871 let mut nn = numbered_network();
872 nn.set_output_activation(ActivationFunction::Tanh);
873
874 nn.set_parameters(&[0.0; 9]);
875
876 assert_eq!(nn.output_activation(), ActivationFunction::Tanh);
877 }
878
879 #[test]
880 fn parameters_round_trip_through_from_parameters() {
881 use rand::SeedableRng;
882 let mut rng = rand::rngs::StdRng::seed_from_u64(7);
883 let original = NeuralNetwork::<3, 2>::new_with_rng(&[4, 5], &mut rng);
884
885 let rebuilt = NeuralNetwork::<3, 2>::from_parameters(
886 &original.hidden_layer_sizes(),
887 &original.parameters(),
888 );
889
890 assert_eq!(rebuilt, original);
891 assert_all_close(
892 &rebuilt.feed_forward(&[0.1, -0.2, 0.3]),
893 &original.feed_forward(&[0.1, -0.2, 0.3]),
894 );
895 }
896
897 #[test]
898 #[should_panic(expected = "expected 9 parameters, got 8")]
899 fn set_parameters_rejects_a_short_list() {
900 numbered_network().set_parameters(&[0.0; 8]);
901 }
902
903 #[test]
904 #[should_panic(expected = "expected 9 parameters, got 10")]
905 fn from_parameters_rejects_a_long_list() {
906 NeuralNetwork::<2, 1>::from_parameters(&[2], &[0.0; 10]);
907 }
908
909 #[test]
910 #[should_panic(expected = "at least one neuron")]
911 fn parameter_count_for_rejects_a_zero_sized_hidden_layer() {
912 NeuralNetwork::<2, 1>::parameter_count_for(&[3, 0]);
913 }
914
915 #[test]
916 fn hidden_layer_sizes_lists_only_hidden_layers() {
917 assert_eq!(NeuralNetwork::<2, 1>::new(&[]).hidden_layer_sizes(), Vec::<usize>::new());
918 assert_eq!(NeuralNetwork::<2, 1>::new(&[5, 3]).hidden_layer_sizes(), vec![5, 3]);
919 }
920
921 #[test]
923 fn networks_can_be_shared_across_threads() {
924 fn assert_send_sync<T: Send + Sync>() {}
925 assert_send_sync::<NeuralNetwork<10, 3>>();
926 }
927}