use crate::{ActivationFunction, Layer, TrainingAlgorithm};
use num_traits::Float;
use rand::distributions::Uniform;
use rand::Rng;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use thiserror::Error;
#[derive(Error, Debug)]
pub enum NetworkError {
#[error("Input size mismatch: expected {expected}, got {actual}")]
InputSizeMismatch { expected: usize, actual: usize },
#[error("Weight count mismatch: expected {expected}, got {actual}")]
WeightCountMismatch { expected: usize, actual: usize },
#[error("Invalid layer configuration")]
InvalidLayerConfiguration,
#[error("Network has no layers")]
NoLayers,
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct Network<T: Float> {
pub layers: Vec<Layer<T>>,
pub connection_rate: T,
}
impl<T: Float> Network<T> {
pub fn new(layer_sizes: &[usize]) -> Self {
NetworkBuilder::new().layers_from_sizes(layer_sizes).build()
}
pub fn num_layers(&self) -> usize {
self.layers.len()
}
pub fn num_inputs(&self) -> usize {
self.layers
.first()
.map(|l| l.num_regular_neurons())
.unwrap_or(0)
}
pub fn num_outputs(&self) -> usize {
self.layers
.last()
.map(|l| l.num_regular_neurons())
.unwrap_or(0)
}
pub fn total_neurons(&self) -> usize {
self.layers.iter().map(|l| l.size()).sum()
}
pub fn total_connections(&self) -> usize {
self.layers
.iter()
.flat_map(|layer| &layer.neurons)
.map(|neuron| neuron.connections.len())
.sum()
}
pub fn get_total_connections(&self) -> usize {
self.total_connections()
}
pub fn run(&mut self, inputs: &[T]) -> Vec<T> {
if self.layers.is_empty() {
return Vec::new();
}
if self.layers[0].set_inputs(inputs).is_err() {
return Vec::new();
}
for i in 1..self.layers.len() {
let prev_outputs = self.layers[i - 1].get_outputs();
self.layers[i].calculate(&prev_outputs);
}
if let Some(output_layer) = self.layers.last() {
output_layer
.neurons
.iter()
.filter(|n| !n.is_bias)
.map(|n| n.value)
.collect()
} else {
Vec::new()
}
}
pub fn get_weights(&self) -> Vec<T> {
let mut weights = Vec::new();
for layer in &self.layers {
for neuron in &layer.neurons {
for connection in &neuron.connections {
weights.push(connection.weight);
}
}
}
weights
}
pub fn set_weights(&mut self, weights: &[T]) -> Result<(), NetworkError> {
let expected = self.total_connections();
if weights.len() != expected {
return Err(NetworkError::WeightCountMismatch {
expected,
actual: weights.len(),
});
}
let mut weight_idx = 0;
for layer in &mut self.layers {
for neuron in &mut layer.neurons {
for connection in &mut neuron.connections {
connection.weight = weights[weight_idx];
weight_idx += 1;
}
}
}
Ok(())
}
pub fn reset(&mut self) {
for layer in &mut self.layers {
layer.reset();
}
}
pub fn set_activation_function_hidden(&mut self, activation_function: ActivationFunction) {
let num_layers = self.layers.len();
if num_layers > 2 {
for i in 1..num_layers - 1 {
self.layers[i].set_activation_function(activation_function);
}
}
}
pub fn set_activation_function_output(&mut self, activation_function: ActivationFunction) {
if let Some(output_layer) = self.layers.last_mut() {
output_layer.set_activation_function(activation_function);
}
}
pub fn set_activation_steepness_hidden(&mut self, steepness: T) {
let num_layers = self.layers.len();
if num_layers > 2 {
for i in 1..num_layers - 1 {
self.layers[i].set_activation_steepness(steepness);
}
}
}
pub fn set_activation_steepness_output(&mut self, steepness: T) {
if let Some(output_layer) = self.layers.last_mut() {
output_layer.set_activation_steepness(steepness);
}
}
pub fn set_activation_function(
&mut self,
layer: usize,
activation_function: ActivationFunction,
) {
if layer < self.layers.len() {
self.layers[layer].set_activation_function(activation_function);
}
}
pub fn randomize_weights(&mut self, min: T, max: T)
where
T: rand::distributions::uniform::SampleUniform,
{
let mut rng = rand::thread_rng();
let range = Uniform::new(min, max);
for layer in &mut self.layers {
for neuron in &mut layer.neurons {
for connection in &mut neuron.connections {
connection.weight = rng.sample(&range);
}
}
}
}
pub fn set_training_algorithm(&mut self, _algorithm: TrainingAlgorithm) {
}
pub fn train(
&mut self,
inputs: &[Vec<T>],
outputs: &[Vec<T>],
learning_rate: f32,
epochs: usize,
) -> Result<(), NetworkError>
where
T: std::ops::AddAssign + std::ops::SubAssign + std::ops::MulAssign + std::cmp::PartialOrd,
{
if inputs.len() != outputs.len() {
return Err(NetworkError::InvalidLayerConfiguration);
}
let lr = T::from(learning_rate as f64).unwrap_or(T::from(0.1).unwrap_or(T::one()));
for _epoch in 0..epochs {
for (input, target) in inputs.iter().zip(outputs.iter()) {
let layer_outputs = self.forward_pass_with_storage(input);
self.backward_pass(&layer_outputs, target, lr);
}
}
Ok(())
}
fn forward_pass_with_storage(&mut self, input: &[T]) -> Vec<Vec<T>> {
let mut layer_outputs = Vec::with_capacity(self.layers.len());
if !self.layers.is_empty() {
let _ = self.layers[0].set_inputs(input);
layer_outputs.push(self.layers[0].get_outputs());
}
for i in 1..self.layers.len() {
let prev_outputs = layer_outputs[i - 1].clone();
self.layers[i].calculate(&prev_outputs);
layer_outputs.push(self.layers[i].get_outputs());
}
layer_outputs
}
fn backward_pass(&mut self, layer_outputs: &[Vec<T>], target: &[T], learning_rate: T) {
if self.layers.is_empty() {
return;
}
let num_layers = self.layers.len();
let mut layer_errors = vec![Vec::new(); num_layers];
if let Some(output_layer) = self.layers.last() {
let output_idx = num_layers - 1;
let outputs = &layer_outputs[output_idx];
for (i, neuron) in output_layer.neurons.iter().enumerate() {
if !neuron.is_bias && i < target.len() && i < outputs.len() {
let error = target[i] - outputs[i];
let delta = error * neuron.activation_derivative();
layer_errors[output_idx].push(delta);
} else {
layer_errors[output_idx].push(T::zero());
}
}
}
for layer_idx in (1..num_layers - 1).rev() {
let current_layer = &self.layers[layer_idx];
let next_layer = &self.layers[layer_idx + 1];
let next_errors = layer_errors[layer_idx + 1].clone();
let mut current_errors = Vec::new();
for (i, neuron) in current_layer.neurons.iter().enumerate() {
if neuron.is_bias {
current_errors.push(T::zero());
continue;
}
let mut error_sum = T::zero();
for (j, next_neuron) in next_layer.neurons.iter().enumerate() {
if !next_neuron.is_bias && j < next_errors.len() {
for connection in &next_neuron.connections {
if connection.from_neuron == i {
error_sum = error_sum + next_errors[j] * connection.weight;
break;
}
}
}
}
let delta = error_sum * neuron.activation_derivative();
current_errors.push(delta);
}
layer_errors[layer_idx] = current_errors;
}
for layer_idx in 1..num_layers {
let prev_outputs = if layer_idx == 1 {
&layer_outputs[0]
} else {
&layer_outputs[layer_idx - 1]
};
let current_errors = &layer_errors[layer_idx];
let current_layer = &mut self.layers[layer_idx];
for (neuron_idx, neuron) in current_layer.neurons.iter_mut().enumerate() {
if neuron.is_bias || neuron_idx >= current_errors.len() {
continue;
}
let error = current_errors[neuron_idx];
for connection in &mut neuron.connections {
if connection.from_neuron < prev_outputs.len() {
let input_value = prev_outputs[connection.from_neuron];
let weight_delta = learning_rate * error * input_value;
connection.weight = connection.weight + weight_delta;
}
}
}
}
}
pub fn run_batch(&mut self, inputs: &[Vec<T>]) -> Vec<Vec<T>> {
inputs.iter().map(|input| self.run(input)).collect()
}
#[cfg(all(feature = "binary", feature = "serde"))]
pub fn to_bytes(&self) -> Vec<u8>
where
T: serde::Serialize,
Network<T>: serde::Serialize,
{
bincode::serialize(self).unwrap_or_default()
}
#[cfg(feature = "binary")]
#[cfg(not(feature = "serde"))]
pub fn to_bytes(&self) -> Vec<u8> {
Vec::new()
}
#[cfg(all(feature = "binary", feature = "serde"))]
pub fn from_bytes(bytes: &[u8]) -> Result<Self, NetworkError>
where
T: serde::de::DeserializeOwned,
Network<T>: serde::de::DeserializeOwned,
{
bincode::deserialize(bytes).map_err(|_| NetworkError::InvalidLayerConfiguration)
}
#[cfg(feature = "binary")]
#[cfg(not(feature = "serde"))]
pub fn from_bytes(_bytes: &[u8]) -> Result<Self, NetworkError> {
Err(NetworkError::InvalidLayerConfiguration)
}
}
pub struct NetworkBuilder<T: Float> {
layers: Vec<(usize, ActivationFunction, T)>,
connection_rate: T,
}
impl<T: Float> NetworkBuilder<T> {
pub fn new() -> Self {
NetworkBuilder {
layers: Vec::new(),
connection_rate: T::one(),
}
}
pub fn layers_from_sizes(mut self, sizes: &[usize]) -> Self {
if sizes.is_empty() {
return self;
}
self.layers
.push((sizes[0], ActivationFunction::Linear, T::one()));
for &size in &sizes[1..sizes.len() - 1] {
self.layers
.push((size, ActivationFunction::Sigmoid, T::one()));
}
if sizes.len() > 1 {
self.layers.push((
sizes[sizes.len() - 1],
ActivationFunction::Sigmoid,
T::one(),
));
}
self
}
pub fn input_layer(mut self, size: usize) -> Self {
self.layers
.push((size, ActivationFunction::Linear, T::one()));
self
}
pub fn hidden_layer(mut self, size: usize) -> Self {
self.layers
.push((size, ActivationFunction::Sigmoid, T::one()));
self
}
pub fn hidden_layer_with_activation(
mut self,
size: usize,
activation: ActivationFunction,
steepness: T,
) -> Self {
self.layers.push((size, activation, steepness));
self
}
pub fn output_layer(mut self, size: usize) -> Self {
self.layers
.push((size, ActivationFunction::Sigmoid, T::one()));
self
}
pub fn output_layer_with_activation(
mut self,
size: usize,
activation: ActivationFunction,
steepness: T,
) -> Self {
self.layers.push((size, activation, steepness));
self
}
pub fn connection_rate(mut self, rate: T) -> Self {
self.connection_rate = rate;
self
}
pub fn build(self) -> Network<T> {
let mut network_layers = Vec::new();
for (i, &(size, activation, steepness)) in self.layers.iter().enumerate() {
let layer = if i == 0 {
Layer::with_bias(size, activation, steepness)
} else if i == self.layers.len() - 1 {
Layer::new(size, activation, steepness)
} else {
Layer::with_bias(size, activation, steepness)
};
network_layers.push(layer);
}
for i in 0..network_layers.len() - 1 {
let (before, after) = network_layers.split_at_mut(i + 1);
before[i].connect_to(&mut after[0], self.connection_rate);
}
Network {
layers: network_layers,
connection_rate: self.connection_rate,
}
}
}
impl<T: Float> Default for NetworkBuilder<T> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_network_builder() {
let network: Network<f32> = NetworkBuilder::new()
.input_layer(2)
.hidden_layer(3)
.output_layer(1)
.build();
assert_eq!(network.num_layers(), 3);
assert_eq!(network.num_inputs(), 2);
assert_eq!(network.num_outputs(), 1);
}
#[test]
fn test_network_run() {
let mut network: Network<f32> = NetworkBuilder::new()
.input_layer(2)
.hidden_layer(3)
.output_layer(1)
.build();
let inputs = vec![0.5, 0.7];
let outputs = network.run(&inputs);
assert_eq!(outputs.len(), 1);
}
#[test]
fn test_total_neurons() {
let network: Network<f32> = NetworkBuilder::new()
.input_layer(2) .hidden_layer(3) .output_layer(1) .build();
assert_eq!(network.total_neurons(), 8);
}
#[test]
fn test_sparse_network() {
let network: Network<f32> = NetworkBuilder::new()
.input_layer(10)
.hidden_layer(10)
.output_layer(10)
.connection_rate(0.5)
.build();
let connections = network.total_connections();
let max_connections = 11 * 10 + 11 * 10;
assert!(connections < max_connections);
}
}