use crate::activation_functions::ActivationFunction;
use crate::layer::Layer;
use crate::ModelError;
use rand::{rng, Rng};
use std::fmt;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(into = "NetworkData", try_from = "NetworkData")]
pub struct NeuralNetwork<const IN: usize, const OUT: usize> {
layers: Vec<Layer>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub(crate) struct NetworkData {
pub(crate) layers: Vec<LayerData>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub(crate) struct LayerData {
pub(crate) activation: ActivationFunction,
pub(crate) weights: Vec<Vec<f64>>,
pub(crate) biases: Vec<f64>,
}
impl<const IN: usize, const OUT: usize> NeuralNetwork<IN, OUT> {
pub fn new(hidden: &[usize]) -> Self {
Self::new_with_rng(hidden, &mut rng())
}
pub fn new_with_rng<R: Rng>(hidden: &[usize], rng: &mut R) -> Self {
Self::build(hidden, |neurons, inputs| {
Layer::random(neurons, inputs, ActivationFunction::default(), rng)
})
}
pub fn from_parameters(hidden: &[usize], parameters: &[f64]) -> Self {
let mut network = Self::build(hidden, |neurons, inputs| {
Layer::zeros(neurons, inputs, ActivationFunction::default())
});
network.set_parameters(parameters);
network
}
fn build(hidden: &[usize], mut make: impl FnMut(usize, usize) -> Layer) -> Self {
let layers = Self::layer_shapes(hidden)
.map(|(neurons, inputs)| make(neurons, inputs))
.collect();
Self { layers }
}
fn layer_shapes(hidden: &[usize]) -> impl Iterator<Item = (usize, usize)> + '_ {
const { assert!(IN > 0, "a network needs at least one input neuron") };
const { assert!(OUT > 0, "a network needs at least one output neuron") };
assert!(
hidden.iter().all(|&size| size > 0),
"every hidden layer must have at least one neuron, got {hidden:?}"
);
let inputs = std::iter::once(IN).chain(hidden.iter().copied());
let neurons = hidden.iter().copied().chain(std::iter::once(OUT));
neurons.zip(inputs)
}
pub fn feed_forward(&self, inputs: &[f64; IN]) -> [f64; OUT] {
let widest = self.layers.iter().map(Layer::size).fold(IN, usize::max);
let mut scratch = vec![0.0; widest * 2];
let (mut current, mut next) = scratch.split_at_mut(widest);
current[..IN].copy_from_slice(inputs);
let mut width = IN;
for layer in &self.layers {
layer.forward_into(¤t[..width], &mut next[..layer.size()]);
width = layer.size();
std::mem::swap(&mut current, &mut next);
}
<[f64; OUT]>::try_from(¤t[..width])
.expect("output layer width should match the OUT type parameter")
}
pub fn parameter_count(&self) -> usize {
self.layers.iter().map(Layer::parameter_count).sum()
}
pub fn parameter_count_for(hidden: &[usize]) -> usize {
Self::layer_shapes(hidden)
.map(|(neurons, inputs)| Layer::parameter_count_for(neurons, inputs))
.sum()
}
pub fn parameters(&self) -> Vec<f64> {
let mut parameters = Vec::with_capacity(self.parameter_count());
for layer in &self.layers {
layer.extend_parameters(&mut parameters);
}
parameters
}
pub fn set_parameters(&mut self, parameters: &[f64]) {
let expected = self.parameter_count();
assert_eq!(
parameters.len(),
expected,
"Incompatible parameter count: expected {expected} parameters, got {}",
parameters.len()
);
let mut rest = parameters;
for layer in &mut self.layers {
rest = layer.take_parameters(rest);
}
}
pub fn set_layer_weights<R: AsRef<[f64]>>(&mut self, layer: usize, weights: &[R]) {
self.layer_mut(layer).set_weights(weights);
}
pub fn layer_weights(&self, layer: usize) -> Vec<Vec<f64>> {
self.layer(layer).weight_rows()
}
pub fn set_layer_biases(&mut self, layer: usize, biases: &[f64]) {
self.layer_mut(layer).set_biases(biases);
}
pub fn layer_biases(&self, layer: usize) -> &[f64] {
self.layer(layer).biases().as_slice()
}
pub fn set_weight(&mut self, layer: usize, neuron: usize, input: usize, weight: f64) {
self.layer_mut(layer).set_weight(neuron, input, weight);
}
pub fn get_weight(&self, layer: usize, neuron: usize, input: usize) -> f64 {
self.layer(layer).weights()[(neuron, input)]
}
pub fn set_bias(&mut self, layer: usize, neuron: usize, bias: f64) {
self.layer_mut(layer).set_bias(neuron, bias);
}
pub fn get_bias(&self, layer: usize, neuron: usize) -> f64 {
self.layer(layer).biases()[neuron]
}
fn layer(&self, layer: usize) -> &Layer {
if layer == 0 {
panic!("Invalid layer index");
}
&self.layers[layer - 1]
}
pub(crate) fn layers(&self) -> &[Layer] {
&self.layers
}
fn layer_mut(&mut self, layer: usize) -> &mut Layer {
if layer == 0 {
panic!("Invalid layer index");
}
&mut self.layers[layer - 1]
}
pub fn num_layers(&self) -> usize {
self.layers.len() + 1
}
pub fn layer_size(&self, layer: usize) -> usize {
if layer == 0 {
return IN;
}
self.layer(layer).size()
}
pub fn hidden_layer_sizes(&self) -> Vec<usize> {
let hidden = &self.layers[..self.layers.len() - 1];
hidden.iter().map(Layer::size).collect()
}
pub fn layer_activation(&self, layer: usize) -> ActivationFunction {
self.layer(layer).activation()
}
pub fn set_layer_activation(&mut self, layer: usize, activation_function: ActivationFunction) {
self.layer_mut(layer).set_activation(activation_function);
}
pub fn set_activation_function(&mut self, activation_function: ActivationFunction) {
for layer in &mut self.layers {
layer.set_activation(activation_function);
}
}
pub fn set_output_activation(&mut self, activation_function: ActivationFunction) {
self.set_layer_activation(self.num_layers() - 1, activation_function);
}
pub fn output_activation(&self) -> ActivationFunction {
self.layer_activation(self.num_layers() - 1)
}
pub fn print(&self) {
for layer in &self.layers {
println!("{:?} {} {}", layer.activation(), layer.weights(), layer.biases());
}
}
}
impl<const IN: usize, const OUT: usize> From<NeuralNetwork<IN, OUT>> for NetworkData {
fn from(network: NeuralNetwork<IN, OUT>) -> Self {
NetworkData::from(&network)
}
}
impl<const IN: usize, const OUT: usize> From<&NeuralNetwork<IN, OUT>> for NetworkData {
fn from(network: &NeuralNetwork<IN, OUT>) -> Self {
let layers = network
.layers
.iter()
.map(|layer| LayerData {
activation: layer.activation(),
weights: layer.weight_rows(),
biases: layer.biases().iter().copied().collect(),
})
.collect();
NetworkData { layers }
}
}
impl<const IN: usize, const OUT: usize> TryFrom<NetworkData> for NeuralNetwork<IN, OUT> {
type Error = ModelError;
fn try_from(data: NetworkData) -> Result<Self, Self::Error> {
let first = data.layers.first().ok_or(ModelError::EmptyNetwork)?;
if let Some(row) = first.weights.first() {
if row.len() != IN {
return Err(ModelError::DimensionMismatch {
end: "input",
expected: IN,
found: row.len(),
});
}
}
let mut inputs = IN;
let mut layers = Vec::with_capacity(data.layers.len());
for (index, layer) in data.layers.iter().enumerate() {
let layer = Layer::from_parts(&layer.weights, &layer.biases, layer.activation, inputs)
.map_err(|reason| ModelError::InconsistentLayer {
layer: index + 1,
reason,
})?;
inputs = layer.size();
layers.push(layer);
}
if inputs != OUT {
return Err(ModelError::DimensionMismatch {
end: "output",
expected: OUT,
found: inputs,
});
}
Ok(Self { layers })
}
}
impl<const IN: usize, const OUT: usize> fmt::Display for NeuralNetwork<IN, OUT> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(f, "Neural Network")?;
writeln!(f)?;
writeln!(f, "Input Layer Size: {IN}")?;
writeln!(f)?;
for (index, layer) in self.layers.iter().enumerate() {
writeln!(f, "Layer {}: {}", index + 1, layer)?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::activation_functions::{binary_step, identity, relu, sigmoid, tanh};
const EPSILON: f64 = 1e-12;
fn assert_all_close(actual: &[f64], expected: &[f64]) {
assert_eq!(actual.len(), expected.len(), "length mismatch");
for (i, (a, e)) in actual.iter().zip(expected).enumerate() {
assert!((a - e).abs() < EPSILON, "at index {i}: expected {e}, got {a}");
}
}
fn fixed_network() -> NeuralNetwork<2, 1> {
let mut nn = NeuralNetwork::<2, 1>::new(&[]);
nn.set_layer_weights(1, &[[0.5, -0.25]]);
nn.set_layer_biases(1, &[0.1]);
nn
}
#[test]
fn new_reports_layer_count_and_sizes() {
let nn = NeuralNetwork::<3, 2>::new(&[5]);
assert_eq!(nn.num_layers(), 3);
assert_eq!(nn.layer_size(0), 3);
assert_eq!(nn.layer_size(1), 5);
assert_eq!(nn.layer_size(2), 2);
}
#[test]
fn new_without_hidden_layers_wires_input_straight_to_output() {
let nn = NeuralNetwork::<3, 2>::new(&[]);
assert_eq!(nn.num_layers(), 2);
assert_eq!(nn.layer_size(0), 3);
assert_eq!(nn.layer_size(1), 2);
}
#[test]
#[should_panic(expected = "at least one neuron")]
fn new_rejects_a_zero_sized_hidden_layer() {
NeuralNetwork::<2, 1>::new(&[0]);
}
#[test]
fn new_with_rng_is_reproducible_for_the_same_seed() {
use rand::SeedableRng;
let mut first_rng = rand::rngs::StdRng::seed_from_u64(42);
let mut second_rng = rand::rngs::StdRng::seed_from_u64(42);
let first = NeuralNetwork::<3, 2>::new_with_rng(&[4], &mut first_rng);
let second = NeuralNetwork::<3, 2>::new_with_rng(&[4], &mut second_rng);
for layer in 1..first.num_layers() {
for neuron in 0..first.layer_size(layer) {
for input in 0..first.layer_size(layer - 1) {
assert_eq!(
first.get_weight(layer, neuron, input),
second.get_weight(layer, neuron, input),
"weight differs at layer {layer}, neuron {neuron}, input {input}"
);
}
}
}
}
#[test]
fn feed_forward_applies_weights_bias_and_activation() {
let nn = fixed_network();
let output = nn.feed_forward(&[1.0, 2.0]);
assert_all_close(&output, &[sigmoid(0.1)]);
}
#[test]
fn feed_forward_defaults_to_sigmoid() {
let nn = fixed_network();
assert_eq!(nn.layer_activation(1), ActivationFunction::Sigmoid);
assert_all_close(&nn.feed_forward(&[1.0, 2.0]), &[sigmoid(0.1)]);
}
#[test]
fn feed_forward_honours_every_activation_function() {
let cases = [
(ActivationFunction::Sigmoid, sigmoid as fn(f64) -> f64),
(ActivationFunction::Tanh, tanh),
(ActivationFunction::ReLU, relu),
(ActivationFunction::BinaryStep, binary_step),
(ActivationFunction::Identity, identity),
];
for (variant, expected) in cases {
let mut nn = fixed_network();
nn.set_activation_function(variant);
assert_eq!(nn.layer_activation(1), variant);
assert_all_close(&nn.feed_forward(&[1.0, 2.0]), &[expected(0.1)]);
}
}
#[test]
fn binary_step_does_not_panic() {
let mut nn = fixed_network();
nn.set_activation_function(ActivationFunction::BinaryStep);
assert_all_close(&nn.feed_forward(&[1.0, 2.0]), &[1.0]);
assert_all_close(&nn.feed_forward(&[-1.0, 2.0]), &[0.0]);
}
#[test]
fn set_and_get_weight_round_trip() {
let mut nn = NeuralNetwork::<2, 2>::new(&[]);
nn.set_weight(1, 1, 0, 0.75);
assert_eq!(nn.get_weight(1, 1, 0), 0.75);
}
#[test]
#[should_panic(expected = "Invalid layer index")]
fn set_layer_weights_rejects_layer_zero() {
let mut nn = NeuralNetwork::<2, 1>::new(&[]);
nn.set_layer_weights(0, &[[0.5, 0.5]]);
}
#[test]
#[should_panic(expected = "Incompatible weights matrix size")]
fn set_layer_weights_rejects_a_mismatched_matrix() {
let mut nn = NeuralNetwork::<2, 1>::new(&[]);
nn.set_layer_weights(1, &[[0.5, 0.5, 0.5]]);
}
#[test]
#[should_panic(expected = "Incompatible biases vector size")]
fn set_layer_biases_rejects_a_mismatched_vector() {
let mut nn = NeuralNetwork::<2, 1>::new(&[]);
nn.set_layer_biases(1, &[0.1, 0.2]);
}
#[test]
fn display_names_the_activation_function() {
let mut nn = fixed_network();
nn.set_activation_function(ActivationFunction::ReLU);
let rendered = nn.to_string();
assert!(
rendered.contains("Activation Function: ReLU"),
"unexpected output:\n{rendered}"
);
assert!(
!rendered.contains("0x"),
"output leaked a pointer address:\n{rendered}"
);
}
#[test]
fn display_reports_the_input_layer_size() {
let nn = NeuralNetwork::<4, 1>::new(&[2]);
assert!(nn.to_string().contains("Input Layer Size: 4"));
}
#[test]
fn layer_weights_and_biases_round_trip_in_row_major_order() {
let mut nn = NeuralNetwork::<2, 3>::new(&[]);
let weights = [[0.1, 0.2], [0.3, 0.4], [0.5, 0.6]];
nn.set_layer_weights(1, &weights);
nn.set_layer_biases(1, &[0.7, 0.8, 0.9]);
assert_eq!(nn.layer_weights(1), weights.map(Vec::from).to_vec());
assert_eq!(nn.layer_biases(1), &[0.7, 0.8, 0.9]);
assert_eq!(nn.get_weight(1, 2, 0), 0.5);
}
#[test]
fn set_layer_weights_accepts_weights_built_at_runtime() {
let mut nn = NeuralNetwork::<2, 1>::new(&[]);
let weights: Vec<Vec<f64>> = vec![vec![0.5, -0.25]];
nn.set_layer_weights(1, &weights);
assert_eq!(nn.layer_weights(1), weights);
}
#[test]
#[should_panic(expected = "Incompatible weights matrix size")]
fn set_layer_weights_rejects_ragged_rows() {
let mut nn = NeuralNetwork::<2, 2>::new(&[]);
let ragged: Vec<Vec<f64>> = vec![vec![0.1, 0.2], vec![0.3]];
nn.set_layer_weights(1, &ragged);
}
#[test]
#[should_panic(expected = "Incompatible weights matrix size")]
fn set_layer_weights_rejects_the_wrong_number_of_rows() {
let mut nn = NeuralNetwork::<2, 2>::new(&[]);
nn.set_layer_weights(1, &[[0.1, 0.2]]);
}
#[test]
fn set_and_get_bias_round_trip() {
let mut nn = NeuralNetwork::<2, 2>::new(&[]);
nn.set_bias(1, 1, -0.3);
assert_eq!(nn.get_bias(1, 1), -0.3);
assert_eq!(nn.layer_biases(1), &[0.0, -0.3]);
}
#[test]
#[should_panic(expected = "Invalid layer index")]
fn layer_biases_rejects_layer_zero() {
let nn = NeuralNetwork::<2, 1>::new(&[]);
nn.layer_biases(0);
}
fn numbered_network() -> NeuralNetwork<2, 1> {
let mut nn = NeuralNetwork::<2, 1>::new(&[2]);
nn.set_layer_weights(1, &[[1.0, 2.0],
[3.0, 4.0]]);
nn.set_layer_biases(1, &[5.0, 6.0]);
nn.set_layer_weights(2, &[[7.0, 8.0]]);
nn.set_layer_biases(2, &[9.0]);
nn
}
#[test]
fn hidden_and_output_layers_can_use_different_activations() {
let mut nn = NeuralNetwork::<1, 1>::new(&[1]);
nn.set_layer_weights(1, &[[1.0]]);
nn.set_layer_weights(2, &[[1.0]]);
nn.set_activation_function(ActivationFunction::ReLU);
nn.set_output_activation(ActivationFunction::Tanh);
assert_all_close(&nn.feed_forward(&[-2.0]), &[0.0]);
assert_all_close(&nn.feed_forward(&[2.0]), &[tanh(2.0)]);
assert_eq!(nn.layer_activation(1), ActivationFunction::ReLU);
assert_eq!(nn.output_activation(), ActivationFunction::Tanh);
}
#[test]
fn set_layer_activation_changes_only_that_layer() {
let mut nn = NeuralNetwork::<2, 1>::new(&[3, 3]);
nn.set_layer_activation(2, ActivationFunction::Identity);
assert_eq!(nn.layer_activation(1), ActivationFunction::Sigmoid);
assert_eq!(nn.layer_activation(2), ActivationFunction::Identity);
assert_eq!(nn.layer_activation(3), ActivationFunction::Sigmoid);
}
#[test]
#[should_panic(expected = "Invalid layer index")]
fn the_input_layer_has_no_activation() {
NeuralNetwork::<2, 1>::new(&[]).layer_activation(0);
}
#[test]
fn parameters_follow_the_documented_order() {
let nn = numbered_network();
assert_eq!(
nn.parameters(),
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0]
);
}
#[test]
fn parameter_count_matches_the_shape() {
let nn = NeuralNetwork::<10, 3>::new(&[8]);
assert_eq!(nn.parameter_count(), 11 * 8 + 9 * 3);
assert_eq!(nn.parameters().len(), nn.parameter_count());
assert_eq!(NeuralNetwork::<10, 3>::parameter_count_for(&[8]), nn.parameter_count());
assert_eq!(NeuralNetwork::<10, 3>::parameter_count_for(&[]), 11 * 3);
}
#[test]
fn set_parameters_writes_the_documented_order() {
let mut nn = NeuralNetwork::<2, 1>::new(&[2]);
nn.set_parameters(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0]);
assert_eq!(nn, numbered_network());
assert_eq!(nn.get_weight(1, 1, 0), 3.0);
assert_eq!(nn.get_bias(2, 0), 9.0);
}
#[test]
fn set_parameters_keeps_activation_functions() {
let mut nn = numbered_network();
nn.set_output_activation(ActivationFunction::Tanh);
nn.set_parameters(&[0.0; 9]);
assert_eq!(nn.output_activation(), ActivationFunction::Tanh);
}
#[test]
fn parameters_round_trip_through_from_parameters() {
use rand::SeedableRng;
let mut rng = rand::rngs::StdRng::seed_from_u64(7);
let original = NeuralNetwork::<3, 2>::new_with_rng(&[4, 5], &mut rng);
let rebuilt = NeuralNetwork::<3, 2>::from_parameters(
&original.hidden_layer_sizes(),
&original.parameters(),
);
assert_eq!(rebuilt, original);
assert_all_close(
&rebuilt.feed_forward(&[0.1, -0.2, 0.3]),
&original.feed_forward(&[0.1, -0.2, 0.3]),
);
}
#[test]
#[should_panic(expected = "expected 9 parameters, got 8")]
fn set_parameters_rejects_a_short_list() {
numbered_network().set_parameters(&[0.0; 8]);
}
#[test]
#[should_panic(expected = "expected 9 parameters, got 10")]
fn from_parameters_rejects_a_long_list() {
NeuralNetwork::<2, 1>::from_parameters(&[2], &[0.0; 10]);
}
#[test]
#[should_panic(expected = "at least one neuron")]
fn parameter_count_for_rejects_a_zero_sized_hidden_layer() {
NeuralNetwork::<2, 1>::parameter_count_for(&[3, 0]);
}
#[test]
fn hidden_layer_sizes_lists_only_hidden_layers() {
assert_eq!(NeuralNetwork::<2, 1>::new(&[]).hidden_layer_sizes(), Vec::<usize>::new());
assert_eq!(NeuralNetwork::<2, 1>::new(&[5, 3]).hidden_layer_sizes(), vec![5, 3]);
}
#[test]
fn networks_can_be_shared_across_threads() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<NeuralNetwork<10, 3>>();
}
}