#[cfg(test)]
mod tests {
use crate::training::*;
use crate::{ActivationFunction, Network};
fn create_xor_data() -> TrainingData<f32> {
TrainingData {
inputs: vec![
vec![0.0, 0.0],
vec![0.0, 1.0],
vec![1.0, 0.0],
vec![1.0, 1.0],
],
outputs: vec![vec![0.0], vec![1.0], vec![1.0], vec![0.0]],
}
}
fn create_simple_network() -> Network<f32> {
let mut network = Network::new(&[2, 3, 1]);
network.set_activation_function_hidden(ActivationFunction::Sigmoid);
network.set_activation_function_output(ActivationFunction::Sigmoid);
network.randomize_weights(-0.5, 0.5);
network
}
#[test]
fn test_adam_training() {
let mut network = create_simple_network();
let data = create_xor_data();
let mut trainer = Adam::new(0.01);
let error = trainer.train_epoch(&mut network, &data).unwrap();
assert!(error.is_finite());
println!("Adam - Initial error: {error}");
}
#[test]
fn test_adamw_training() {
let mut network = create_simple_network();
let data = create_xor_data();
let mut trainer = AdamW::new(0.01).with_weight_decay(0.001);
let error = trainer.train_epoch(&mut network, &data).unwrap();
assert!(error.is_finite());
println!("AdamW - Initial error: {error}");
}
#[test]
fn test_incremental_backprop_training() {
let mut network = create_simple_network();
let data = create_xor_data();
let mut trainer = IncrementalBackprop::new(0.1).with_momentum(0.9);
let error = trainer.train_epoch(&mut network, &data).unwrap();
assert!(error.is_finite());
println!("IncrementalBackprop - Initial error: {error}");
}
#[test]
fn test_batch_backprop_training() {
let mut network = create_simple_network();
let data = create_xor_data();
let mut trainer = BatchBackprop::new(0.1).with_momentum(0.9);
let error = trainer.train_epoch(&mut network, &data).unwrap();
assert!(error.is_finite());
println!("BatchBackprop - Initial error: {error}");
}
#[test]
fn test_rprop_training() {
let mut network = create_simple_network();
let data = create_xor_data();
let mut trainer = Rprop::new();
let error = trainer.train_epoch(&mut network, &data).unwrap();
assert!(error.is_finite());
println!("Rprop - Initial error: {error}");
}
#[test]
fn test_quickprop_training() {
let mut network = create_simple_network();
let data = create_xor_data();
let mut trainer = Quickprop::new();
let error = trainer.train_epoch(&mut network, &data).unwrap();
assert!(error.is_finite());
println!("Quickprop - Initial error: {error}");
}
#[test]
fn test_all_algorithms_improve_error() {
let data = create_xor_data();
let algorithms: Vec<(&str, Box<dyn TrainingAlgorithm<f32>>)> = vec![
("Adam", Box::new(Adam::new(0.1))), ("AdamW", Box::new(AdamW::new(0.1))),
(
"IncrementalBackprop",
Box::new(IncrementalBackprop::new(0.1)),
),
("BatchBackprop", Box::new(BatchBackprop::new(0.1))),
("Rprop", Box::new(Rprop::new())),
("Quickprop", Box::new(Quickprop::new())),
];
for (name, mut trainer) in algorithms {
let mut network = create_simple_network();
let initial_error = trainer.calculate_error(&network, &data);
let mut min_error = initial_error;
for epoch in 0..50 {
let error = trainer.train_epoch(&mut network, &data).unwrap();
if error < min_error {
min_error = error;
}
if epoch % 10 == 0 {
println!("{name} - Epoch {epoch}: error = {error:.6}");
}
}
let final_error = trainer.calculate_error(&network, &data);
println!("{}: Initial error: {:.6}, Final error: {:.6}, Min error: {:.6}, Improvement: {:.2}%",
name, initial_error, final_error, min_error,
(1.0 - min_error/initial_error) * 100.0);
assert!(
min_error <= initial_error * 1.1,
"{name} error increased significantly. Initial: {initial_error}, Min: {min_error}"
);
}
}
}