Skip to main content

only_brain/
neural_network.rs

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/// Neural Network
10///
11/// This is the main struct of the library: a chain of fully connected layers, each with
12/// its own weights, biases and [`ActivationFunction`]. You can use this struct and its
13/// methods to create, manipulate and even implement your own ways to train a neural
14/// network.
15///
16/// The number of input neurons (`IN`) and output neurons (`OUT`) are part of the type,
17/// so feeding a wrongly sized input is a compile error rather than a runtime panic. The
18/// hidden layers stay dynamic and are given at construction time.
19///
20/// # Layers
21///
22/// Layer 0 is the input layer, which only passes the inputs on, so it has no weights,
23/// biases or activation. Layers `1..num_layers()` each hold one row of weights per
24/// neuron, with one weight per neuron of the previous layer, and one bias per neuron.
25///
26/// # Example
27///
28/// ```
29/// use only_brain::NeuralNetwork;
30///
31/// // A 2 -> 2 -> 1 network.
32/// let mut nn = NeuralNetwork::<2, 1>::new(&[2]);
33///
34/// // One row per neuron, one weight per neuron of the previous layer.
35/// nn.set_layer_weights(1, &[[0.1, 0.2],
36///                           [0.3, 0.4]]);
37/// nn.set_layer_biases(1, &[0.1, 0.2]);
38///
39/// nn.set_layer_weights(2, &[[0.9, 0.8]]);
40/// nn.set_layer_biases(2, &[0.1]);
41///
42/// let output = nn.feed_forward(&[0.5, 0.2]);
43///
44/// println!("{:?}", output);
45/// ```
46///
47/// # Flat parameter view
48///
49/// Search methods such as genetic algorithms usually work on a flat list of numbers
50/// rather than on layers. [`parameters`](Self::parameters),
51/// [`set_parameters`](Self::set_parameters) and
52/// [`from_parameters`](Self::from_parameters) convert between a network and such a
53/// list, in an order that is documented and stable across versions:
54///
55/// - layer by layer, from layer 1 to the output layer;
56/// - within a layer, every weight first, row by row (all the weights of neuron 0, then
57///   of neuron 1, and so on, each row in the order of the previous layer's neurons);
58/// - then that layer's biases, one per neuron.
59///
60/// Activation functions are part of the network's shape, not of its parameters, so
61/// they are left untouched by [`set_parameters`](Self::set_parameters).
62///
63/// ```
64/// use only_brain::NeuralNetwork;
65///
66/// let mut nn = NeuralNetwork::<2, 1>::new(&[]);
67/// nn.set_layer_weights(1, &[[0.5, -0.25]]);
68/// nn.set_layer_biases(1, &[0.1]);
69///
70/// assert_eq!(nn.parameters(), vec![0.5, -0.25, 0.1]);
71///
72/// let rebuilt = NeuralNetwork::<2, 1>::from_parameters(&[], &nn.parameters());
73/// assert_eq!(rebuilt, nn);
74/// ```
75///
76/// # Serialization
77///
78/// Networks implement serde's `Serialize` and `Deserialize` through a plain form, a
79/// list of layers each with its `activation`, `weights` (one row per neuron) and
80/// `biases`, so they can be embedded in your own types and formats. Deserializing checks
81/// that the layers fit together and match `IN` and `OUT`. To save a network to a file,
82/// see [`crate::dump_model`] and [`crate::load_model`].
83#[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/// The plain form of a network, before its shape has been checked.
90///
91/// Serde and the model file readers both go through this, so every entry point, not
92/// only [`crate::load_model`], rejects layers that do not fit together. It is also what
93/// keeps the serialized form independent of how layers are stored in memory.
94#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
95pub(crate) struct NetworkData {
96    pub(crate) layers: Vec<LayerData>,
97}
98
99/// The plain form of one layer.
100#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
101pub(crate) struct LayerData {
102    pub(crate) activation: ActivationFunction,
103    /// One row per neuron, one weight per neuron of the previous layer.
104    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    /// Creates a new Neural Network with the given hidden layer sizes. The input and
110    /// output widths come from the type parameters, so `hidden` lists only the layers
111    /// between them and may be empty.
112    ///
113    /// Weights are initialised uniformly at random in `[-1, 1)`, biases at zero, and
114    /// every layer uses [`ActivationFunction::Sigmoid`]. Use
115    /// [`NeuralNetwork::new_with_rng`] to control the seed.
116    ///
117    /// # Panics
118    ///
119    /// Panics if any hidden layer size is zero. `IN` or `OUT` being zero is a compile
120    /// error.
121    ///
122    /// # Example
123    ///
124    /// ```
125    /// # use only_brain::NeuralNetwork;
126    /// // 2 -> 2 -> 1
127    /// let nn = NeuralNetwork::<2, 1>::new(&[2]);
128    ///
129    /// // 3 -> 1, no hidden layers
130    /// let direct = NeuralNetwork::<3, 1>::new(&[]);
131    /// ```
132    pub fn new(hidden: &[usize]) -> Self {
133        Self::new_with_rng(hidden, &mut rng())
134    }
135
136    /// Creates a new Neural Network using the given random number generator.
137    ///
138    /// Seeding the generator makes initialisation reproducible, which is what you want
139    /// in tests and when a training run needs to be repeatable.
140    ///
141    /// # Panics
142    ///
143    /// Panics if any hidden layer size is zero.
144    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    /// Creates a network with the given hidden layer sizes from a flat list of
151    /// parameters, in the order described in the [flat parameter
152    /// view](Self#flat-parameter-view). Every layer uses
153    /// [`ActivationFunction::Sigmoid`]; change that afterwards with
154    /// [`set_activation_function`](Self::set_activation_function) and friends.
155    ///
156    /// # Panics
157    ///
158    /// Panics if any hidden layer size is zero, or if `parameters` does not hold exactly
159    /// [`parameter_count_for(hidden)`](Self::parameter_count_for) values.
160    ///
161    /// # Example
162    ///
163    /// ```
164    /// # use only_brain::{ActivationFunction, NeuralNetwork};
165    /// // A genome from a genetic algorithm, for a 2 -> 2 -> 1 network.
166    /// let genome = vec![0.5; NeuralNetwork::<2, 1>::parameter_count_for(&[2])];
167    ///
168    /// let mut nn = NeuralNetwork::<2, 1>::from_parameters(&[2], &genome);
169    /// nn.set_output_activation(ActivationFunction::Tanh);
170    ///
171    /// let [steering] = nn.feed_forward(&[0.3, -0.8]);
172    /// assert!((-1.0..=1.0).contains(&steering));
173    /// ```
174    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    /// Chains `IN`, the hidden sizes and `OUT` into layers made by `make(neurons, inputs)`.
183    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    /// The `(neurons, inputs)` of every layer after the input layer.
192    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    /// Feeds the given inputs to the neural network and returns the output.
206    ///
207    /// The input and output widths are checked at compile time.
208    ///
209    /// # Example
210    ///
211    /// ```
212    /// # use only_brain::NeuralNetwork;
213    /// let mut nn = NeuralNetwork::<1, 1>::new(&[]);
214    ///
215    /// nn.set_layer_weights(1, &[[0.5]]);
216    /// nn.set_layer_biases(1, &[0.5]);
217    ///
218    /// let output = nn.feed_forward(&[0.5]);
219    /// assert!((output[0] - 0.679178699175393).abs() < 1e-12);
220    /// ```
221    pub fn feed_forward(&self, inputs: &[f64; IN]) -> [f64; OUT] {
222        // Two halves that the layers ping-pong between, each wide enough for the
223        // widest layer. Networks up to STACK_WIDTH neurons per layer, which is every
224        // network a genome is likely to drive, run without touching the heap.
225        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    /// The width of the widest layer, counting the input layer.
238    fn widest_layer(&self) -> usize {
239        self.layers.iter().map(Layer::size).fold(IN, usize::max)
240    }
241
242    /// Runs the layers over `scratch`, which must hold `2 * widest_layer()` values.
243    #[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(&current[..width], &mut next[..layer.size()]);
252            width = layer.size();
253            std::mem::swap(&mut current, &mut next);
254        }
255
256        // The layer chain is built from IN..OUT and every setter checks its dimensions,
257        // so the final slice always has exactly OUT elements.
258        <[f64; OUT]>::try_from(&current[..width])
259            .expect("output layer width should match the OUT type parameter")
260    }
261
262    /// Feeds every input to the network and returns one output per input, in order.
263    ///
264    /// This computes the same values as calling [`feed_forward`](Self::feed_forward)
265    /// once per input, but runs each layer as a single matrix product over all the
266    /// inputs at once, which is faster when there are many of them: scoring a dataset,
267    /// for example. It does not help when every input goes to a different network.
268    ///
269    /// # Example
270    ///
271    /// ```
272    /// # use only_brain::NeuralNetwork;
273    /// let nn = NeuralNetwork::<2, 1>::new(&[3]);
274    ///
275    /// let inputs = [[0.0, 0.0], [0.0, 1.0], [1.0, 0.0], [1.0, 1.0]];
276    /// let outputs = nn.feed_forward_batch(&inputs);
277    ///
278    /// assert_eq!(outputs.len(), 4);
279    /// assert_eq!(outputs[2], nn.feed_forward(&inputs[2]));
280    /// ```
281    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        // One column per sample, with as many rows as the widest layer; each layer
286        // works on the top rows it needs and the two matrices ping-pong.
287        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            // As in `feed_forward`, the last layer is always OUT wide.
305            .map(|column| std::array::from_fn(|neuron| column[neuron]))
306            .collect()
307    }
308
309    /// Returns the number of weights and biases in the network, which is the length of
310    /// [`parameters`](Self::parameters).
311    pub fn parameter_count(&self) -> usize {
312        self.layers.iter().map(Layer::parameter_count).sum()
313    }
314
315    /// Returns the number of weights and biases a network with these hidden layer sizes
316    /// has, without building one. This is the genome length to give a genetic algorithm.
317    ///
318    /// # Panics
319    ///
320    /// Panics if any hidden layer size is zero.
321    ///
322    /// # Example
323    ///
324    /// ```
325    /// # use only_brain::NeuralNetwork;
326    /// // 10 -> 8 -> 3: (10 + 1) * 8 + (8 + 1) * 3
327    /// assert_eq!(NeuralNetwork::<10, 3>::parameter_count_for(&[8]), 115);
328    /// ```
329    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    /// Returns every weight and bias as one flat list, in the order described in the
336    /// [flat parameter view](Self#flat-parameter-view).
337    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    /// Replaces every weight and bias from one flat list, in the order described in the
346    /// [flat parameter view](Self#flat-parameter-view). Activation functions are kept.
347    ///
348    /// # Panics
349    ///
350    /// Panics if `parameters` does not hold exactly
351    /// [`parameter_count()`](Self::parameter_count) values.
352    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    /// Sets every weight of the given layer, as one row per neuron of that layer
368    /// holding one weight per neuron of the previous layer.
369    ///
370    /// Layer 0 is the input layer, which has no weights, so `layer` starts at 1. Nested
371    /// arrays and `Vec<Vec<f64>>` are both accepted.
372    ///
373    /// # Panics
374    ///
375    /// Panics if `layer` is 0 or past the output layer, or if `weights` does not have
376    /// exactly `layer_size(layer)` rows of `layer_size(layer - 1)` weights.
377    ///
378    /// # Example
379    ///
380    /// ```
381    /// # use only_brain::NeuralNetwork;
382    /// let mut nn = NeuralNetwork::<3, 1>::new(&[2]);
383    ///
384    /// nn.set_layer_weights(1, &[[0.1, 0.2, 0.3],
385    ///                           [0.4, 0.5, 0.6]]);
386    ///
387    /// // Weights computed at runtime work the same way.
388    /// let output_weights = vec![vec![0.7, 0.8]];
389    /// nn.set_layer_weights(2, &output_weights);
390    ///
391    /// assert_eq!(nn.layer_weights(2), output_weights);
392    /// ```
393    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    /// Returns the weights of the given layer, as one row per neuron of that layer
398    /// holding one weight per neuron of the previous layer.
399    ///
400    /// # Panics
401    ///
402    /// Panics if `layer` is 0 or past the output layer.
403    pub fn layer_weights(&self, layer: usize) -> Vec<Vec<f64>> {
404        self.layer(layer).weight_rows()
405    }
406
407    /// Sets the biases of the given layer, one per neuron.
408    ///
409    /// # Panics
410    ///
411    /// Panics if `layer` is 0 or past the output layer, or if `biases` does not have
412    /// exactly `layer_size(layer)` elements.
413    pub fn set_layer_biases(&mut self, layer: usize, biases: &[f64]) {
414        self.layer_mut(layer).set_biases(biases);
415    }
416
417    /// Returns the biases of the given layer, one per neuron.
418    ///
419    /// # Panics
420    ///
421    /// Panics if `layer` is 0 or past the output layer.
422    pub fn layer_biases(&self, layer: usize) -> &[f64] {
423        self.layer(layer).biases().as_slice()
424    }
425
426    /// Sets the weight of a specific neuron connection. The layer index must be greater
427    /// than 0 since the input layer does not have weights.
428    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    /// Gets the weight of a specific neuron connection. The layer index must be greater
433    /// than 0 since the input layer does not have weights.
434    pub fn get_weight(&self, layer: usize, neuron: usize, input: usize) -> f64 {
435        self.layer(layer).weights()[(neuron, input)]
436    }
437
438    /// Sets the bias of a specific neuron. The layer index must be greater than 0 since
439    /// the input layer does not have biases.
440    pub fn set_bias(&mut self, layer: usize, neuron: usize, bias: f64) {
441        self.layer_mut(layer).set_bias(neuron, bias);
442    }
443
444    /// Gets the bias of a specific neuron. The layer index must be greater than 0 since
445    /// the input layer does not have biases.
446    pub fn get_bias(&self, layer: usize, neuron: usize) -> f64 {
447        self.layer(layer).biases()[neuron]
448    }
449
450    /// The layer that receives weights for `layer`, which counts the input layer as 0.
451    fn layer(&self, layer: usize) -> &Layer {
452        if layer == 0 {
453            panic!("Invalid layer index");
454        }
455        &self.layers[layer - 1]
456    }
457
458    /// Every layer after the input layer, for the model writer.
459    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    /// Returns the number of layers of the neural network, counting the input layer.
471    pub fn num_layers(&self) -> usize {
472        self.layers.len() + 1
473    }
474
475    /// Returns the number of neurons of the given layer. Layer 0 is the input layer.
476    pub fn layer_size(&self, layer: usize) -> usize {
477        if layer == 0 {
478            return IN;
479        }
480        self.layer(layer).size()
481    }
482
483    /// Returns the sizes of the hidden layers, the same list given to
484    /// [`new`](Self::new) or [`from_parameters`](Self::from_parameters).
485    ///
486    /// ```
487    /// # use only_brain::NeuralNetwork;
488    /// let nn = NeuralNetwork::<4, 2>::new(&[8, 6]);
489    ///
490    /// let copy = NeuralNetwork::<4, 2>::from_parameters(&nn.hidden_layer_sizes(), &nn.parameters());
491    /// assert_eq!(copy.hidden_layer_sizes(), vec![8, 6]);
492    /// ```
493    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    /// Returns the activation function of the given layer.
499    ///
500    /// # Panics
501    ///
502    /// Panics if `layer` is 0, since the input layer has no activation, or past the
503    /// output layer.
504    pub fn layer_activation(&self, layer: usize) -> ActivationFunction {
505        self.layer(layer).activation()
506    }
507
508    /// Sets the activation function of the given layer.
509    ///
510    /// # Panics
511    ///
512    /// Panics if `layer` is 0, since the input layer has no activation, or past the
513    /// output layer.
514    pub fn set_layer_activation(&mut self, layer: usize, activation_function: ActivationFunction) {
515        self.layer_mut(layer).set_activation(activation_function);
516    }
517
518    /// Sets the activation function of every layer of the network.
519    ///
520    /// Combine it with [`set_output_activation`](Self::set_output_activation) to give
521    /// the hidden layers and the output layer different functions.
522    ///
523    /// # Example
524    ///
525    /// ```
526    /// # use only_brain::{ActivationFunction, NeuralNetwork};
527    /// let mut nn = NeuralNetwork::<2, 1>::new(&[3]);
528    /// nn.set_activation_function(ActivationFunction::ReLU);
529    /// nn.set_output_activation(ActivationFunction::Tanh);
530    ///
531    /// assert_eq!(nn.layer_activation(1), ActivationFunction::ReLU);
532    /// assert_eq!(nn.layer_activation(2), ActivationFunction::Tanh);
533    /// ```
534    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    /// Sets the activation function of the output layer only.
541    pub fn set_output_activation(&mut self, activation_function: ActivationFunction) {
542        self.set_layer_activation(self.num_layers() - 1, activation_function);
543    }
544
545    /// Returns the activation function of the output layer.
546    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    /// Checks that the stored layers chain from `IN` inputs to `OUT` outputs and that
583    /// each layer agrees with itself.
584    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    /// A 2 -> 1 network with known weights, so outputs can be checked by hand.
650    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    /// A network with no hidden layers is still a valid two-layer network. The old
668    /// "fewer than two layers" runtime check is gone because IN and OUT guarantee it.
669    #[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        // 0.5 * 1.0 + (-0.25) * 2.0 + 0.1 = 0.1
712        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    /// The activation function used to be a field with no setter, so every
762    /// network silently ran sigmoid regardless of what was configured.
763    #[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    /// `BinaryStep` was missing from the old lookup table and panicked here.
783    #[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    /// `Display` used to print the address of a function pointer here.
822    #[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        // Row = neuron, column = input, matching set_weight and get_weight.
857        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    /// A 2 -> 2 -> 1 network whose every parameter is distinct, so a wrong order shows.
903    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        // relu(-2) = 0, then tanh(0) = 0; relu(2) = 2, then tanh(2).
922        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    /// Fitness functions run in parallel (rayon), which needs Send + Sync.
1027    #[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}