1use crate::activation_functions::ActivationFunction;
2use crate::layer::Layer;
3use crate::ModelError;
4use nalgebra::DMatrix;
5use rand::{rng, Rng};
6use std::fmt;
7use serde::{Deserialize, Serialize};
8
9#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
84#[serde(into = "NetworkData", try_from = "NetworkData")]
85pub struct NeuralNetwork<const IN: usize, const OUT: usize> {
86 layers: Vec<Layer>,
87}
88
89#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
95pub(crate) struct NetworkData {
96 pub(crate) layers: Vec<LayerData>,
97}
98
99#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
101pub(crate) struct LayerData {
102 pub(crate) activation: ActivationFunction,
103 pub(crate) weights: Vec<Vec<f64>>,
105 pub(crate) biases: Vec<f64>,
106}
107
108impl<const IN: usize, const OUT: usize> NeuralNetwork<IN, OUT> {
109 pub fn new(hidden: &[usize]) -> Self {
133 Self::new_with_rng(hidden, &mut rng())
134 }
135
136 pub fn new_with_rng<R: Rng>(hidden: &[usize], rng: &mut R) -> Self {
145 Self::build(hidden, |neurons, inputs| {
146 Layer::random(neurons, inputs, ActivationFunction::default(), rng)
147 })
148 }
149
150 pub fn from_parameters(hidden: &[usize], parameters: &[f64]) -> Self {
175 let mut network = Self::build(hidden, |neurons, inputs| {
176 Layer::zeros(neurons, inputs, ActivationFunction::default())
177 });
178 network.set_parameters(parameters);
179 network
180 }
181
182 fn build(hidden: &[usize], mut make: impl FnMut(usize, usize) -> Layer) -> Self {
184 let layers = Self::layer_shapes(hidden)
185 .map(|(neurons, inputs)| make(neurons, inputs))
186 .collect();
187
188 Self { layers }
189 }
190
191 fn layer_shapes(hidden: &[usize]) -> impl Iterator<Item = (usize, usize)> + '_ {
193 const { assert!(IN > 0, "a network needs at least one input neuron") };
194 const { assert!(OUT > 0, "a network needs at least one output neuron") };
195 assert!(
196 hidden.iter().all(|&size| size > 0),
197 "every hidden layer must have at least one neuron, got {hidden:?}"
198 );
199
200 let inputs = std::iter::once(IN).chain(hidden.iter().copied());
201 let neurons = hidden.iter().copied().chain(std::iter::once(OUT));
202 neurons.zip(inputs)
203 }
204
205 pub fn feed_forward(&self, inputs: &[f64; IN]) -> [f64; OUT] {
222 const STACK_WIDTH: usize = 16;
226
227 let widest = self.widest_layer();
228 if widest <= STACK_WIDTH {
229 let mut scratch = [0.0; STACK_WIDTH * 2];
230 self.forward_in(inputs, &mut scratch[..widest * 2])
231 } else {
232 let mut scratch = vec![0.0; widest * 2];
233 self.forward_in(inputs, &mut scratch)
234 }
235 }
236
237 fn widest_layer(&self) -> usize {
239 self.layers.iter().map(Layer::size).fold(IN, usize::max)
240 }
241
242 #[inline(always)]
244 fn forward_in(&self, inputs: &[f64; IN], scratch: &mut [f64]) -> [f64; OUT] {
245 let (mut current, mut next) = scratch.split_at_mut(scratch.len() / 2);
246
247 current[..IN].copy_from_slice(inputs);
248 let mut width = IN;
249
250 for layer in &self.layers {
251 layer.forward_into(¤t[..width], &mut next[..layer.size()]);
252 width = layer.size();
253 std::mem::swap(&mut current, &mut next);
254 }
255
256 <[f64; OUT]>::try_from(¤t[..width])
259 .expect("output layer width should match the OUT type parameter")
260 }
261
262 pub fn feed_forward_batch(&self, inputs: &[[f64; IN]]) -> Vec<[f64; OUT]> {
282 let samples = inputs.len();
283 let widest = self.widest_layer();
284
285 let mut current = DMatrix::zeros(widest, samples);
288 let mut next = DMatrix::zeros(widest, samples);
289
290 for (mut column, sample) in current.column_iter_mut().zip(inputs) {
291 column.rows_mut(0, IN).copy_from_slice(sample);
292 }
293 let mut width = IN;
294
295 for layer in &self.layers {
296 layer.forward_batch_into(current.rows(0, width), next.rows_mut(0, layer.size()));
297 width = layer.size();
298 std::mem::swap(&mut current, &mut next);
299 }
300
301 current
302 .rows(0, width)
303 .column_iter()
304 .map(|column| std::array::from_fn(|neuron| column[neuron]))
306 .collect()
307 }
308
309 pub fn parameter_count(&self) -> usize {
312 self.layers.iter().map(Layer::parameter_count).sum()
313 }
314
315 pub fn parameter_count_for(hidden: &[usize]) -> usize {
330 Self::layer_shapes(hidden)
331 .map(|(neurons, inputs)| Layer::parameter_count_for(neurons, inputs))
332 .sum()
333 }
334
335 pub fn parameters(&self) -> Vec<f64> {
338 let mut parameters = Vec::with_capacity(self.parameter_count());
339 for layer in &self.layers {
340 layer.extend_parameters(&mut parameters);
341 }
342 parameters
343 }
344
345 pub fn set_parameters(&mut self, parameters: &[f64]) {
353 let expected = self.parameter_count();
354 assert_eq!(
355 parameters.len(),
356 expected,
357 "Incompatible parameter count: expected {expected} parameters, got {}",
358 parameters.len()
359 );
360
361 let mut rest = parameters;
362 for layer in &mut self.layers {
363 rest = layer.take_parameters(rest);
364 }
365 }
366
367 pub fn set_layer_weights<R: AsRef<[f64]>>(&mut self, layer: usize, weights: &[R]) {
394 self.layer_mut(layer).set_weights(weights);
395 }
396
397 pub fn layer_weights(&self, layer: usize) -> Vec<Vec<f64>> {
404 self.layer(layer).weight_rows()
405 }
406
407 pub fn set_layer_biases(&mut self, layer: usize, biases: &[f64]) {
414 self.layer_mut(layer).set_biases(biases);
415 }
416
417 pub fn layer_biases(&self, layer: usize) -> &[f64] {
423 self.layer(layer).biases().as_slice()
424 }
425
426 pub fn set_weight(&mut self, layer: usize, neuron: usize, input: usize, weight: f64) {
429 self.layer_mut(layer).set_weight(neuron, input, weight);
430 }
431
432 pub fn get_weight(&self, layer: usize, neuron: usize, input: usize) -> f64 {
435 self.layer(layer).weights()[(neuron, input)]
436 }
437
438 pub fn set_bias(&mut self, layer: usize, neuron: usize, bias: f64) {
441 self.layer_mut(layer).set_bias(neuron, bias);
442 }
443
444 pub fn get_bias(&self, layer: usize, neuron: usize) -> f64 {
447 self.layer(layer).biases()[neuron]
448 }
449
450 fn layer(&self, layer: usize) -> &Layer {
452 if layer == 0 {
453 panic!("Invalid layer index");
454 }
455 &self.layers[layer - 1]
456 }
457
458 pub(crate) fn layers(&self) -> &[Layer] {
460 &self.layers
461 }
462
463 fn layer_mut(&mut self, layer: usize) -> &mut Layer {
464 if layer == 0 {
465 panic!("Invalid layer index");
466 }
467 &mut self.layers[layer - 1]
468 }
469
470 pub fn num_layers(&self) -> usize {
472 self.layers.len() + 1
473 }
474
475 pub fn layer_size(&self, layer: usize) -> usize {
477 if layer == 0 {
478 return IN;
479 }
480 self.layer(layer).size()
481 }
482
483 pub fn hidden_layer_sizes(&self) -> Vec<usize> {
494 let hidden = &self.layers[..self.layers.len() - 1];
495 hidden.iter().map(Layer::size).collect()
496 }
497
498 pub fn layer_activation(&self, layer: usize) -> ActivationFunction {
505 self.layer(layer).activation()
506 }
507
508 pub fn set_layer_activation(&mut self, layer: usize, activation_function: ActivationFunction) {
515 self.layer_mut(layer).set_activation(activation_function);
516 }
517
518 pub fn set_activation_function(&mut self, activation_function: ActivationFunction) {
535 for layer in &mut self.layers {
536 layer.set_activation(activation_function);
537 }
538 }
539
540 pub fn set_output_activation(&mut self, activation_function: ActivationFunction) {
542 self.set_layer_activation(self.num_layers() - 1, activation_function);
543 }
544
545 pub fn output_activation(&self) -> ActivationFunction {
547 self.layer_activation(self.num_layers() - 1)
548 }
549
550 pub fn print(&self) {
551 for layer in &self.layers {
552 println!("{:?} {} {}", layer.activation(), layer.weights(), layer.biases());
553 }
554 }
555}
556
557impl<const IN: usize, const OUT: usize> From<NeuralNetwork<IN, OUT>> for NetworkData {
558 fn from(network: NeuralNetwork<IN, OUT>) -> Self {
559 NetworkData::from(&network)
560 }
561}
562
563impl<const IN: usize, const OUT: usize> From<&NeuralNetwork<IN, OUT>> for NetworkData {
564 fn from(network: &NeuralNetwork<IN, OUT>) -> Self {
565 let layers = network
566 .layers
567 .iter()
568 .map(|layer| LayerData {
569 activation: layer.activation(),
570 weights: layer.weight_rows(),
571 biases: layer.biases().iter().copied().collect(),
572 })
573 .collect();
574
575 NetworkData { layers }
576 }
577}
578
579impl<const IN: usize, const OUT: usize> TryFrom<NetworkData> for NeuralNetwork<IN, OUT> {
580 type Error = ModelError;
581
582 fn try_from(data: NetworkData) -> Result<Self, Self::Error> {
585 let first = data.layers.first().ok_or(ModelError::EmptyNetwork)?;
586 if let Some(row) = first.weights.first() {
587 if row.len() != IN {
588 return Err(ModelError::DimensionMismatch {
589 end: "input",
590 expected: IN,
591 found: row.len(),
592 });
593 }
594 }
595
596 let mut inputs = IN;
597 let mut layers = Vec::with_capacity(data.layers.len());
598 for (index, layer) in data.layers.iter().enumerate() {
599 let layer = Layer::from_parts(&layer.weights, &layer.biases, layer.activation, inputs)
600 .map_err(|reason| ModelError::InconsistentLayer {
601 layer: index + 1,
602 reason,
603 })?;
604 inputs = layer.size();
605 layers.push(layer);
606 }
607
608 if inputs != OUT {
609 return Err(ModelError::DimensionMismatch {
610 end: "output",
611 expected: OUT,
612 found: inputs,
613 });
614 }
615
616 Ok(Self { layers })
617 }
618}
619
620impl<const IN: usize, const OUT: usize> fmt::Display for NeuralNetwork<IN, OUT> {
621 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
622 writeln!(f, "Neural Network")?;
623 writeln!(f)?;
624 writeln!(f, "Input Layer Size: {IN}")?;
625 writeln!(f)?;
626 for (index, layer) in self.layers.iter().enumerate() {
627 writeln!(f, "Layer {}: {}", index + 1, layer)?;
628 }
629 Ok(())
630 }
631}
632
633#[cfg(test)]
634mod tests {
635 use super::*;
636 use crate::activation_functions::{binary_step, identity, relu, sigmoid, tanh};
637 use rand::rngs::StdRng;
638 use rand::SeedableRng;
639
640 const EPSILON: f64 = 1e-12;
641
642 fn assert_all_close(actual: &[f64], expected: &[f64]) {
643 assert_eq!(actual.len(), expected.len(), "length mismatch");
644 for (i, (a, e)) in actual.iter().zip(expected).enumerate() {
645 assert!((a - e).abs() < EPSILON, "at index {i}: expected {e}, got {a}");
646 }
647 }
648
649 fn fixed_network() -> NeuralNetwork<2, 1> {
651 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
652 nn.set_layer_weights(1, &[[0.5, -0.25]]);
653 nn.set_layer_biases(1, &[0.1]);
654 nn
655 }
656
657 #[test]
658 fn new_reports_layer_count_and_sizes() {
659 let nn = NeuralNetwork::<3, 2>::new(&[5]);
660
661 assert_eq!(nn.num_layers(), 3);
662 assert_eq!(nn.layer_size(0), 3);
663 assert_eq!(nn.layer_size(1), 5);
664 assert_eq!(nn.layer_size(2), 2);
665 }
666
667 #[test]
670 fn new_without_hidden_layers_wires_input_straight_to_output() {
671 let nn = NeuralNetwork::<3, 2>::new(&[]);
672
673 assert_eq!(nn.num_layers(), 2);
674 assert_eq!(nn.layer_size(0), 3);
675 assert_eq!(nn.layer_size(1), 2);
676 }
677
678 #[test]
679 #[should_panic(expected = "at least one neuron")]
680 fn new_rejects_a_zero_sized_hidden_layer() {
681 NeuralNetwork::<2, 1>::new(&[0]);
682 }
683
684 #[test]
685 fn new_with_rng_is_reproducible_for_the_same_seed() {
686 use rand::SeedableRng;
687
688 let mut first_rng = rand::rngs::StdRng::seed_from_u64(42);
689 let mut second_rng = rand::rngs::StdRng::seed_from_u64(42);
690
691 let first = NeuralNetwork::<3, 2>::new_with_rng(&[4], &mut first_rng);
692 let second = NeuralNetwork::<3, 2>::new_with_rng(&[4], &mut second_rng);
693
694 for layer in 1..first.num_layers() {
695 for neuron in 0..first.layer_size(layer) {
696 for input in 0..first.layer_size(layer - 1) {
697 assert_eq!(
698 first.get_weight(layer, neuron, input),
699 second.get_weight(layer, neuron, input),
700 "weight differs at layer {layer}, neuron {neuron}, input {input}"
701 );
702 }
703 }
704 }
705 }
706
707 #[test]
708 fn feed_forward_applies_weights_bias_and_activation() {
709 let nn = fixed_network();
710
711 let output = nn.feed_forward(&[1.0, 2.0]);
713
714 assert_all_close(&output, &[sigmoid(0.1)]);
715 }
716
717 #[test]
718 fn feed_forward_batch_matches_feed_forward_for_every_input() {
719 let mut rng = StdRng::seed_from_u64(7);
720 let mut nn = NeuralNetwork::<3, 2>::new_with_rng(&[5, 4], &mut rng);
721 nn.set_layer_activation(1, ActivationFunction::ReLU);
722 nn.set_layer_activation(2, ActivationFunction::Tanh);
723
724 let inputs: Vec<[f64; 3]> = (0..7).map(|i| [i as f64 * 0.3 - 1.0, 0.5, -(i as f64)]).collect();
725
726 let batch = nn.feed_forward_batch(&inputs);
727
728 assert_eq!(batch.len(), inputs.len());
729 for (input, output) in inputs.iter().zip(&batch) {
730 assert_all_close(output, &nn.feed_forward(input));
731 }
732 }
733
734 #[test]
735 fn feed_forward_batch_of_nothing_is_empty() {
736 let nn = NeuralNetwork::<2, 1>::new(&[2]);
737
738 assert!(nn.feed_forward_batch(&[]).is_empty());
739 }
740
741 #[test]
742 fn feed_forward_batch_works_without_hidden_layers() {
743 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
744 nn.set_layer_weights(1, &[[0.5, -0.25]]);
745 nn.set_layer_biases(1, &[0.1]);
746 nn.set_output_activation(ActivationFunction::Identity);
747
748 let outputs = nn.feed_forward_batch(&[[1.0, 1.0], [2.0, 0.0]]);
749
750 assert_all_close(&outputs.concat(), &[0.35, 1.1]);
751 }
752
753 #[test]
754 fn feed_forward_defaults_to_sigmoid() {
755 let nn = fixed_network();
756
757 assert_eq!(nn.layer_activation(1), ActivationFunction::Sigmoid);
758 assert_all_close(&nn.feed_forward(&[1.0, 2.0]), &[sigmoid(0.1)]);
759 }
760
761 #[test]
764 fn feed_forward_honours_every_activation_function() {
765 let cases = [
766 (ActivationFunction::Sigmoid, sigmoid as fn(f64) -> f64),
767 (ActivationFunction::Tanh, tanh),
768 (ActivationFunction::ReLU, relu),
769 (ActivationFunction::BinaryStep, binary_step),
770 (ActivationFunction::Identity, identity),
771 ];
772
773 for (variant, expected) in cases {
774 let mut nn = fixed_network();
775 nn.set_activation_function(variant);
776
777 assert_eq!(nn.layer_activation(1), variant);
778 assert_all_close(&nn.feed_forward(&[1.0, 2.0]), &[expected(0.1)]);
779 }
780 }
781
782 #[test]
784 fn binary_step_does_not_panic() {
785 let mut nn = fixed_network();
786 nn.set_activation_function(ActivationFunction::BinaryStep);
787
788 assert_all_close(&nn.feed_forward(&[1.0, 2.0]), &[1.0]);
789 assert_all_close(&nn.feed_forward(&[-1.0, 2.0]), &[0.0]);
790 }
791
792 #[test]
793 fn set_and_get_weight_round_trip() {
794 let mut nn = NeuralNetwork::<2, 2>::new(&[]);
795 nn.set_weight(1, 1, 0, 0.75);
796
797 assert_eq!(nn.get_weight(1, 1, 0), 0.75);
798 }
799
800 #[test]
801 #[should_panic(expected = "Invalid layer index")]
802 fn set_layer_weights_rejects_layer_zero() {
803 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
804 nn.set_layer_weights(0, &[[0.5, 0.5]]);
805 }
806
807 #[test]
808 #[should_panic(expected = "Incompatible weights matrix size")]
809 fn set_layer_weights_rejects_a_mismatched_matrix() {
810 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
811 nn.set_layer_weights(1, &[[0.5, 0.5, 0.5]]);
812 }
813
814 #[test]
815 #[should_panic(expected = "Incompatible biases vector size")]
816 fn set_layer_biases_rejects_a_mismatched_vector() {
817 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
818 nn.set_layer_biases(1, &[0.1, 0.2]);
819 }
820
821 #[test]
823 fn display_names_the_activation_function() {
824 let mut nn = fixed_network();
825 nn.set_activation_function(ActivationFunction::ReLU);
826
827 let rendered = nn.to_string();
828
829 assert!(
830 rendered.contains("Activation Function: ReLU"),
831 "unexpected output:\n{rendered}"
832 );
833 assert!(
834 !rendered.contains("0x"),
835 "output leaked a pointer address:\n{rendered}"
836 );
837 }
838
839 #[test]
840 fn display_reports_the_input_layer_size() {
841 let nn = NeuralNetwork::<4, 1>::new(&[2]);
842
843 assert!(nn.to_string().contains("Input Layer Size: 4"));
844 }
845
846 #[test]
847 fn layer_weights_and_biases_round_trip_in_row_major_order() {
848 let mut nn = NeuralNetwork::<2, 3>::new(&[]);
849 let weights = [[0.1, 0.2], [0.3, 0.4], [0.5, 0.6]];
850
851 nn.set_layer_weights(1, &weights);
852 nn.set_layer_biases(1, &[0.7, 0.8, 0.9]);
853
854 assert_eq!(nn.layer_weights(1), weights.map(Vec::from).to_vec());
855 assert_eq!(nn.layer_biases(1), &[0.7, 0.8, 0.9]);
856 assert_eq!(nn.get_weight(1, 2, 0), 0.5);
858 }
859
860 #[test]
861 fn set_layer_weights_accepts_weights_built_at_runtime() {
862 let mut nn = NeuralNetwork::<2, 1>::new(&[]);
863 let weights: Vec<Vec<f64>> = vec![vec![0.5, -0.25]];
864
865 nn.set_layer_weights(1, &weights);
866
867 assert_eq!(nn.layer_weights(1), weights);
868 }
869
870 #[test]
871 #[should_panic(expected = "Incompatible weights matrix size")]
872 fn set_layer_weights_rejects_ragged_rows() {
873 let mut nn = NeuralNetwork::<2, 2>::new(&[]);
874 let ragged: Vec<Vec<f64>> = vec![vec![0.1, 0.2], vec![0.3]];
875
876 nn.set_layer_weights(1, &ragged);
877 }
878
879 #[test]
880 #[should_panic(expected = "Incompatible weights matrix size")]
881 fn set_layer_weights_rejects_the_wrong_number_of_rows() {
882 let mut nn = NeuralNetwork::<2, 2>::new(&[]);
883 nn.set_layer_weights(1, &[[0.1, 0.2]]);
884 }
885
886 #[test]
887 fn set_and_get_bias_round_trip() {
888 let mut nn = NeuralNetwork::<2, 2>::new(&[]);
889 nn.set_bias(1, 1, -0.3);
890
891 assert_eq!(nn.get_bias(1, 1), -0.3);
892 assert_eq!(nn.layer_biases(1), &[0.0, -0.3]);
893 }
894
895 #[test]
896 #[should_panic(expected = "Invalid layer index")]
897 fn layer_biases_rejects_layer_zero() {
898 let nn = NeuralNetwork::<2, 1>::new(&[]);
899 nn.layer_biases(0);
900 }
901
902 fn numbered_network() -> NeuralNetwork<2, 1> {
904 let mut nn = NeuralNetwork::<2, 1>::new(&[2]);
905 nn.set_layer_weights(1, &[[1.0, 2.0],
906 [3.0, 4.0]]);
907 nn.set_layer_biases(1, &[5.0, 6.0]);
908 nn.set_layer_weights(2, &[[7.0, 8.0]]);
909 nn.set_layer_biases(2, &[9.0]);
910 nn
911 }
912
913 #[test]
914 fn hidden_and_output_layers_can_use_different_activations() {
915 let mut nn = NeuralNetwork::<1, 1>::new(&[1]);
916 nn.set_layer_weights(1, &[[1.0]]);
917 nn.set_layer_weights(2, &[[1.0]]);
918 nn.set_activation_function(ActivationFunction::ReLU);
919 nn.set_output_activation(ActivationFunction::Tanh);
920
921 assert_all_close(&nn.feed_forward(&[-2.0]), &[0.0]);
923 assert_all_close(&nn.feed_forward(&[2.0]), &[tanh(2.0)]);
924 assert_eq!(nn.layer_activation(1), ActivationFunction::ReLU);
925 assert_eq!(nn.output_activation(), ActivationFunction::Tanh);
926 }
927
928 #[test]
929 fn set_layer_activation_changes_only_that_layer() {
930 let mut nn = NeuralNetwork::<2, 1>::new(&[3, 3]);
931 nn.set_layer_activation(2, ActivationFunction::Identity);
932
933 assert_eq!(nn.layer_activation(1), ActivationFunction::Sigmoid);
934 assert_eq!(nn.layer_activation(2), ActivationFunction::Identity);
935 assert_eq!(nn.layer_activation(3), ActivationFunction::Sigmoid);
936 }
937
938 #[test]
939 #[should_panic(expected = "Invalid layer index")]
940 fn the_input_layer_has_no_activation() {
941 NeuralNetwork::<2, 1>::new(&[]).layer_activation(0);
942 }
943
944 #[test]
945 fn parameters_follow_the_documented_order() {
946 let nn = numbered_network();
947
948 assert_eq!(
949 nn.parameters(),
950 vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0]
951 );
952 }
953
954 #[test]
955 fn parameter_count_matches_the_shape() {
956 let nn = NeuralNetwork::<10, 3>::new(&[8]);
957
958 assert_eq!(nn.parameter_count(), 11 * 8 + 9 * 3);
959 assert_eq!(nn.parameters().len(), nn.parameter_count());
960 assert_eq!(NeuralNetwork::<10, 3>::parameter_count_for(&[8]), nn.parameter_count());
961 assert_eq!(NeuralNetwork::<10, 3>::parameter_count_for(&[]), 11 * 3);
962 }
963
964 #[test]
965 fn set_parameters_writes_the_documented_order() {
966 let mut nn = NeuralNetwork::<2, 1>::new(&[2]);
967 nn.set_parameters(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0]);
968
969 assert_eq!(nn, numbered_network());
970 assert_eq!(nn.get_weight(1, 1, 0), 3.0);
971 assert_eq!(nn.get_bias(2, 0), 9.0);
972 }
973
974 #[test]
975 fn set_parameters_keeps_activation_functions() {
976 let mut nn = numbered_network();
977 nn.set_output_activation(ActivationFunction::Tanh);
978
979 nn.set_parameters(&[0.0; 9]);
980
981 assert_eq!(nn.output_activation(), ActivationFunction::Tanh);
982 }
983
984 #[test]
985 fn parameters_round_trip_through_from_parameters() {
986 use rand::SeedableRng;
987 let mut rng = rand::rngs::StdRng::seed_from_u64(7);
988 let original = NeuralNetwork::<3, 2>::new_with_rng(&[4, 5], &mut rng);
989
990 let rebuilt = NeuralNetwork::<3, 2>::from_parameters(
991 &original.hidden_layer_sizes(),
992 &original.parameters(),
993 );
994
995 assert_eq!(rebuilt, original);
996 assert_all_close(
997 &rebuilt.feed_forward(&[0.1, -0.2, 0.3]),
998 &original.feed_forward(&[0.1, -0.2, 0.3]),
999 );
1000 }
1001
1002 #[test]
1003 #[should_panic(expected = "expected 9 parameters, got 8")]
1004 fn set_parameters_rejects_a_short_list() {
1005 numbered_network().set_parameters(&[0.0; 8]);
1006 }
1007
1008 #[test]
1009 #[should_panic(expected = "expected 9 parameters, got 10")]
1010 fn from_parameters_rejects_a_long_list() {
1011 NeuralNetwork::<2, 1>::from_parameters(&[2], &[0.0; 10]);
1012 }
1013
1014 #[test]
1015 #[should_panic(expected = "at least one neuron")]
1016 fn parameter_count_for_rejects_a_zero_sized_hidden_layer() {
1017 NeuralNetwork::<2, 1>::parameter_count_for(&[3, 0]);
1018 }
1019
1020 #[test]
1021 fn hidden_layer_sizes_lists_only_hidden_layers() {
1022 assert_eq!(NeuralNetwork::<2, 1>::new(&[]).hidden_layer_sizes(), Vec::<usize>::new());
1023 assert_eq!(NeuralNetwork::<2, 1>::new(&[5, 3]).hidden_layer_sizes(), vec![5, 3]);
1024 }
1025
1026 #[test]
1028 fn networks_can_be_shared_across_threads() {
1029 fn assert_send_sync<T: Send + Sync>() {}
1030 assert_send_sync::<NeuralNetwork<10, 3>>();
1031 }
1032}