Skip to main content

only_brain/
neural_network.rs

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/// Neural Network
9///
10/// This is the main struct of the library: a chain of fully connected layers, each with
11/// its own weights, biases and [`ActivationFunction`]. You can use this struct and its
12/// methods to create, manipulate and even implement your own ways to train a neural
13/// network.
14///
15/// The number of input neurons (`IN`) and output neurons (`OUT`) are part of the type,
16/// so feeding a wrongly sized input is a compile error rather than a runtime panic. The
17/// hidden layers stay dynamic and are given at construction time.
18///
19/// # Layers
20///
21/// Layer 0 is the input layer, which only passes the inputs on, so it has no weights,
22/// biases or activation. Layers `1..num_layers()` each hold one row of weights per
23/// neuron, with one weight per neuron of the previous layer, and one bias per neuron.
24///
25/// # Example
26///
27/// ```
28/// use only_brain::NeuralNetwork;
29///
30/// // A 2 -> 2 -> 1 network.
31/// let mut nn = NeuralNetwork::<2, 1>::new(&[2]);
32///
33/// // One row per neuron, one weight per neuron of the previous layer.
34/// nn.set_layer_weights(1, &[[0.1, 0.2],
35///                           [0.3, 0.4]]);
36/// nn.set_layer_biases(1, &[0.1, 0.2]);
37///
38/// nn.set_layer_weights(2, &[[0.9, 0.8]]);
39/// nn.set_layer_biases(2, &[0.1]);
40///
41/// let output = nn.feed_forward(&[0.5, 0.2]);
42///
43/// println!("{:?}", output);
44/// ```
45///
46/// # Flat parameter view
47///
48/// Search methods such as genetic algorithms usually work on a flat list of numbers
49/// rather than on layers. [`parameters`](Self::parameters),
50/// [`set_parameters`](Self::set_parameters) and
51/// [`from_parameters`](Self::from_parameters) convert between a network and such a
52/// list, in an order that is documented and stable across versions:
53///
54/// - layer by layer, from layer 1 to the output layer;
55/// - within a layer, every weight first, row by row (all the weights of neuron 0, then
56///   of neuron 1, and so on, each row in the order of the previous layer's neurons);
57/// - then that layer's biases, one per neuron.
58///
59/// Activation functions are part of the network's shape, not of its parameters, so
60/// they are left untouched by [`set_parameters`](Self::set_parameters).
61///
62/// ```
63/// use only_brain::NeuralNetwork;
64///
65/// let mut nn = NeuralNetwork::<2, 1>::new(&[]);
66/// nn.set_layer_weights(1, &[[0.5, -0.25]]);
67/// nn.set_layer_biases(1, &[0.1]);
68///
69/// assert_eq!(nn.parameters(), vec![0.5, -0.25, 0.1]);
70///
71/// let rebuilt = NeuralNetwork::<2, 1>::from_parameters(&[], &nn.parameters());
72/// assert_eq!(rebuilt, nn);
73/// ```
74///
75/// # Serialization
76///
77/// Networks implement serde's `Serialize` and `Deserialize` through a plain form, a
78/// list of layers each with its `activation`, `weights` (one row per neuron) and
79/// `biases`, so they can be embedded in your own types and formats. Deserializing checks
80/// that the layers fit together and match `IN` and `OUT`. To save a network to a file,
81/// see [`crate::dump_model`] and [`crate::load_model`].
82#[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/// The plain form of a network, before its shape has been checked.
89///
90/// Serde and the model file readers both go through this, so every entry point, not
91/// only [`crate::load_model`], rejects layers that do not fit together. It is also what
92/// keeps the serialized form independent of how layers are stored in memory.
93#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
94pub(crate) struct NetworkData {
95    pub(crate) layers: Vec<LayerData>,
96}
97
98/// The plain form of one layer.
99#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
100pub(crate) struct LayerData {
101    pub(crate) activation: ActivationFunction,
102    /// One row per neuron, one weight per neuron of the previous layer.
103    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    /// Creates a new Neural Network with the given hidden layer sizes. The input and
109    /// output widths come from the type parameters, so `hidden` lists only the layers
110    /// between them and may be empty.
111    ///
112    /// Weights are initialised uniformly at random in `[-1, 1)`, biases at zero, and
113    /// every layer uses [`ActivationFunction::Sigmoid`]. Use
114    /// [`NeuralNetwork::new_with_rng`] to control the seed.
115    ///
116    /// # Panics
117    ///
118    /// Panics if any hidden layer size is zero. `IN` or `OUT` being zero is a compile
119    /// error.
120    ///
121    /// # Example
122    ///
123    /// ```
124    /// # use only_brain::NeuralNetwork;
125    /// // 2 -> 2 -> 1
126    /// let nn = NeuralNetwork::<2, 1>::new(&[2]);
127    ///
128    /// // 3 -> 1, no hidden layers
129    /// let direct = NeuralNetwork::<3, 1>::new(&[]);
130    /// ```
131    pub fn new(hidden: &[usize]) -> Self {
132        Self::new_with_rng(hidden, &mut rng())
133    }
134
135    /// Creates a new Neural Network using the given random number generator.
136    ///
137    /// Seeding the generator makes initialisation reproducible, which is what you want
138    /// in tests and when a training run needs to be repeatable.
139    ///
140    /// # Panics
141    ///
142    /// Panics if any hidden layer size is zero.
143    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    /// Creates a network with the given hidden layer sizes from a flat list of
150    /// parameters, in the order described in the [flat parameter
151    /// view](Self#flat-parameter-view). Every layer uses
152    /// [`ActivationFunction::Sigmoid`]; change that afterwards with
153    /// [`set_activation_function`](Self::set_activation_function) and friends.
154    ///
155    /// # Panics
156    ///
157    /// Panics if any hidden layer size is zero, or if `parameters` does not hold exactly
158    /// [`parameter_count_for(hidden)`](Self::parameter_count_for) values.
159    ///
160    /// # Example
161    ///
162    /// ```
163    /// # use only_brain::{ActivationFunction, NeuralNetwork};
164    /// // A genome from a genetic algorithm, for a 2 -> 2 -> 1 network.
165    /// let genome = vec![0.5; NeuralNetwork::<2, 1>::parameter_count_for(&[2])];
166    ///
167    /// let mut nn = NeuralNetwork::<2, 1>::from_parameters(&[2], &genome);
168    /// nn.set_output_activation(ActivationFunction::Tanh);
169    ///
170    /// let [steering] = nn.feed_forward(&[0.3, -0.8]);
171    /// assert!((-1.0..=1.0).contains(&steering));
172    /// ```
173    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    /// Chains `IN`, the hidden sizes and `OUT` into layers made by `make(neurons, inputs)`.
182    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    /// The `(neurons, inputs)` of every layer after the input layer.
191    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    /// Feeds the given inputs to the neural network and returns the output.
205    ///
206    /// The input and output widths are checked at compile time.
207    ///
208    /// # Example
209    ///
210    /// ```
211    /// # use only_brain::NeuralNetwork;
212    /// let mut nn = NeuralNetwork::<1, 1>::new(&[]);
213    ///
214    /// nn.set_layer_weights(1, &[[0.5]]);
215    /// nn.set_layer_biases(1, &[0.5]);
216    ///
217    /// let output = nn.feed_forward(&[0.5]);
218    /// assert!((output[0] - 0.679178699175393).abs() < 1e-12);
219    /// ```
220    pub fn feed_forward(&self, inputs: &[f64; IN]) -> [f64; OUT] {
221        // One scratch allocation split into two halves that the layers ping-pong
222        // between, each wide enough for the widest layer.
223        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(&current[..width], &mut next[..layer.size()]);
232            width = layer.size();
233            std::mem::swap(&mut current, &mut next);
234        }
235
236        // The layer chain is built from IN..OUT and every setter checks its dimensions,
237        // so the final slice always has exactly OUT elements.
238        <[f64; OUT]>::try_from(&current[..width])
239            .expect("output layer width should match the OUT type parameter")
240    }
241
242    /// Returns the number of weights and biases in the network, which is the length of
243    /// [`parameters`](Self::parameters).
244    pub fn parameter_count(&self) -> usize {
245        self.layers.iter().map(Layer::parameter_count).sum()
246    }
247
248    /// Returns the number of weights and biases a network with these hidden layer sizes
249    /// has, without building one. This is the genome length to give a genetic algorithm.
250    ///
251    /// # Panics
252    ///
253    /// Panics if any hidden layer size is zero.
254    ///
255    /// # Example
256    ///
257    /// ```
258    /// # use only_brain::NeuralNetwork;
259    /// // 10 -> 8 -> 3: (10 + 1) * 8 + (8 + 1) * 3
260    /// assert_eq!(NeuralNetwork::<10, 3>::parameter_count_for(&[8]), 115);
261    /// ```
262    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    /// Returns every weight and bias as one flat list, in the order described in the
269    /// [flat parameter view](Self#flat-parameter-view).
270    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    /// Replaces every weight and bias from one flat list, in the order described in the
279    /// [flat parameter view](Self#flat-parameter-view). Activation functions are kept.
280    ///
281    /// # Panics
282    ///
283    /// Panics if `parameters` does not hold exactly
284    /// [`parameter_count()`](Self::parameter_count) values.
285    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    /// Sets every weight of the given layer, as one row per neuron of that layer
301    /// holding one weight per neuron of the previous layer.
302    ///
303    /// Layer 0 is the input layer, which has no weights, so `layer` starts at 1. Nested
304    /// arrays and `Vec<Vec<f64>>` are both accepted.
305    ///
306    /// # Panics
307    ///
308    /// Panics if `layer` is 0 or past the output layer, or if `weights` does not have
309    /// exactly `layer_size(layer)` rows of `layer_size(layer - 1)` weights.
310    ///
311    /// # Example
312    ///
313    /// ```
314    /// # use only_brain::NeuralNetwork;
315    /// let mut nn = NeuralNetwork::<3, 1>::new(&[2]);
316    ///
317    /// nn.set_layer_weights(1, &[[0.1, 0.2, 0.3],
318    ///                           [0.4, 0.5, 0.6]]);
319    ///
320    /// // Weights computed at runtime work the same way.
321    /// let output_weights = vec![vec![0.7, 0.8]];
322    /// nn.set_layer_weights(2, &output_weights);
323    ///
324    /// assert_eq!(nn.layer_weights(2), output_weights);
325    /// ```
326    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    /// Returns the weights of the given layer, as one row per neuron of that layer
331    /// holding one weight per neuron of the previous layer.
332    ///
333    /// # Panics
334    ///
335    /// Panics if `layer` is 0 or past the output layer.
336    pub fn layer_weights(&self, layer: usize) -> Vec<Vec<f64>> {
337        self.layer(layer).weight_rows()
338    }
339
340    /// Sets the biases of the given layer, one per neuron.
341    ///
342    /// # Panics
343    ///
344    /// Panics if `layer` is 0 or past the output layer, or if `biases` does not have
345    /// exactly `layer_size(layer)` elements.
346    pub fn set_layer_biases(&mut self, layer: usize, biases: &[f64]) {
347        self.layer_mut(layer).set_biases(biases);
348    }
349
350    /// Returns the biases of the given layer, one per neuron.
351    ///
352    /// # Panics
353    ///
354    /// Panics if `layer` is 0 or past the output layer.
355    pub fn layer_biases(&self, layer: usize) -> &[f64] {
356        self.layer(layer).biases().as_slice()
357    }
358
359    /// Sets the weight of a specific neuron connection. The layer index must be greater
360    /// than 0 since the input layer does not have weights.
361    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    /// Gets the weight of a specific neuron connection. The layer index must be greater
366    /// than 0 since the input layer does not have weights.
367    pub fn get_weight(&self, layer: usize, neuron: usize, input: usize) -> f64 {
368        self.layer(layer).weights()[(neuron, input)]
369    }
370
371    /// Sets the bias of a specific neuron. The layer index must be greater than 0 since
372    /// the input layer does not have biases.
373    pub fn set_bias(&mut self, layer: usize, neuron: usize, bias: f64) {
374        self.layer_mut(layer).set_bias(neuron, bias);
375    }
376
377    /// Gets the bias of a specific neuron. The layer index must be greater than 0 since
378    /// the input layer does not have biases.
379    pub fn get_bias(&self, layer: usize, neuron: usize) -> f64 {
380        self.layer(layer).biases()[neuron]
381    }
382
383    /// The layer that receives weights for `layer`, which counts the input layer as 0.
384    fn layer(&self, layer: usize) -> &Layer {
385        if layer == 0 {
386            panic!("Invalid layer index");
387        }
388        &self.layers[layer - 1]
389    }
390
391    /// Every layer after the input layer, for the model writer.
392    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    /// Returns the number of layers of the neural network, counting the input layer.
404    pub fn num_layers(&self) -> usize {
405        self.layers.len() + 1
406    }
407
408    /// Returns the number of neurons of the given layer. Layer 0 is the input layer.
409    pub fn layer_size(&self, layer: usize) -> usize {
410        if layer == 0 {
411            return IN;
412        }
413        self.layer(layer).size()
414    }
415
416    /// Returns the sizes of the hidden layers, the same list given to
417    /// [`new`](Self::new) or [`from_parameters`](Self::from_parameters).
418    ///
419    /// ```
420    /// # use only_brain::NeuralNetwork;
421    /// let nn = NeuralNetwork::<4, 2>::new(&[8, 6]);
422    ///
423    /// let copy = NeuralNetwork::<4, 2>::from_parameters(&nn.hidden_layer_sizes(), &nn.parameters());
424    /// assert_eq!(copy.hidden_layer_sizes(), vec![8, 6]);
425    /// ```
426    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    /// Returns the activation function of the given layer.
432    ///
433    /// # Panics
434    ///
435    /// Panics if `layer` is 0, since the input layer has no activation, or past the
436    /// output layer.
437    pub fn layer_activation(&self, layer: usize) -> ActivationFunction {
438        self.layer(layer).activation()
439    }
440
441    /// Sets the activation function of the given layer.
442    ///
443    /// # Panics
444    ///
445    /// Panics if `layer` is 0, since the input layer has no activation, or past the
446    /// output layer.
447    pub fn set_layer_activation(&mut self, layer: usize, activation_function: ActivationFunction) {
448        self.layer_mut(layer).set_activation(activation_function);
449    }
450
451    /// Sets the activation function of every layer of the network.
452    ///
453    /// Combine it with [`set_output_activation`](Self::set_output_activation) to give
454    /// the hidden layers and the output layer different functions.
455    ///
456    /// # Example
457    ///
458    /// ```
459    /// # use only_brain::{ActivationFunction, NeuralNetwork};
460    /// let mut nn = NeuralNetwork::<2, 1>::new(&[3]);
461    /// nn.set_activation_function(ActivationFunction::ReLU);
462    /// nn.set_output_activation(ActivationFunction::Tanh);
463    ///
464    /// assert_eq!(nn.layer_activation(1), ActivationFunction::ReLU);
465    /// assert_eq!(nn.layer_activation(2), ActivationFunction::Tanh);
466    /// ```
467    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    /// Sets the activation function of the output layer only.
474    pub fn set_output_activation(&mut self, activation_function: ActivationFunction) {
475        self.set_layer_activation(self.num_layers() - 1, activation_function);
476    }
477
478    /// Returns the activation function of the output layer.
479    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    /// Checks that the stored layers chain from `IN` inputs to `OUT` outputs and that
516    /// each layer agrees with itself.
517    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    /// A 2 -> 1 network with known weights, so outputs can be checked by hand.
581    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    /// A network with no hidden layers is still a valid two-layer network. The old
599    /// "fewer than two layers" runtime check is gone because IN and OUT guarantee it.
600    #[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        // 0.5 * 1.0 + (-0.25) * 2.0 + 0.1 = 0.1
643        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    /// The activation function used to be a field with no setter, so every
657    /// network silently ran sigmoid regardless of what was configured.
658    #[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    /// `BinaryStep` was missing from the old lookup table and panicked here.
678    #[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    /// `Display` used to print the address of a function pointer here.
717    #[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        // Row = neuron, column = input, matching set_weight and get_weight.
752        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    /// A 2 -> 2 -> 1 network whose every parameter is distinct, so a wrong order shows.
798    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        // relu(-2) = 0, then tanh(0) = 0; relu(2) = 2, then tanh(2).
817        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    /// Fitness functions run in parallel (rayon), which needs Send + Sync.
922    #[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}