use num_traits::Float;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
use thiserror::Error;
use crate::{
cascade_error,
errors::{CascadeErrorCategory, RuvFannError},
ActivationFunction, Network, TrainingData,
};
#[cfg(feature = "logging")]
use log::{debug, error, info};
#[derive(Error, Debug)]
pub enum CascadeError {
#[error("Candidate generation failed: {0}")]
CandidateGeneration(String),
#[error("Candidate training failed: {0}")]
CandidateTraining(String),
#[error("No suitable candidate found")]
NoSuitableCandidate,
#[error("Network topology modification failed: {0}")]
TopologyModification(String),
#[error("Correlation calculation failed: {0}")]
CorrelationCalculation(String),
#[error("Output training failed: {0}")]
OutputTraining(String),
#[error("Invalid cascade configuration: {0}")]
InvalidConfiguration(String),
#[error("Convergence failure: {0}")]
ConvergenceFailure(String),
}
#[derive(Debug, Clone)]
pub struct CascadeNetwork<T: Float> {
pub network: Network<T>,
pub config: CascadeConfig<T>,
pub state: CascadeState<T>,
}
impl<T: Float> CascadeNetwork<T> {
pub fn new(network: Network<T>, config: CascadeConfig<T>) -> Self {
Self {
network,
config,
state: CascadeState::default(),
}
}
}
#[derive(Debug, Clone)]
pub struct CascadeState<T: Float> {
pub hidden_neurons_added: usize,
pub current_error: T,
pub best_correlation: T,
}
impl<T: Float> Default for CascadeState<T> {
fn default() -> Self {
Self {
hidden_neurons_added: 0,
current_error: T::infinity(),
best_correlation: T::zero(),
}
}
}
#[derive(Debug, Clone)]
pub struct CascadeConfig<T: Float> {
pub max_hidden_neurons: usize,
pub num_candidates: usize,
pub output_max_epochs: usize,
pub candidate_max_epochs: usize,
pub output_learning_rate: T,
pub candidate_learning_rate: T,
pub output_target_error: T,
pub candidate_target_correlation: T,
pub min_correlation_improvement: T,
pub candidate_weight_range: (T, T),
pub candidate_activations: Vec<ActivationFunction>,
pub patience: usize,
pub use_weight_decay: bool,
pub weight_decay: T,
pub use_momentum: bool,
pub momentum: T,
pub parallel_candidates: bool,
pub random_seed: Option<u64>,
pub verbose: bool,
}
impl<T: Float> Default for CascadeConfig<T> {
fn default() -> Self {
Self {
max_hidden_neurons: 100,
num_candidates: 8,
output_max_epochs: 1000,
candidate_max_epochs: 1000,
output_learning_rate: T::from(0.1).unwrap(),
candidate_learning_rate: T::from(0.1).unwrap(),
output_target_error: T::from(0.01).unwrap(),
candidate_target_correlation: T::from(0.4).unwrap(),
min_correlation_improvement: T::from(0.01).unwrap(),
candidate_weight_range: (T::from(-1.0).unwrap(), T::from(1.0).unwrap()),
candidate_activations: vec![
ActivationFunction::Sigmoid,
ActivationFunction::Tanh,
ActivationFunction::Gaussian,
],
patience: 50,
use_weight_decay: true,
weight_decay: T::from(0.0001).unwrap(),
use_momentum: true,
momentum: T::from(0.9).unwrap(),
parallel_candidates: true,
random_seed: None,
verbose: false,
}
}
}
#[derive(Debug, Clone)]
pub struct CandidateNeuron<T: Float> {
pub weights: Vec<T>,
pub bias: T,
pub activation: ActivationFunction,
pub steepness: T,
pub correlation: T,
pub training_history: Vec<T>,
pub output: T,
pub weight_gradients: Vec<T>,
pub bias_gradient: T,
pub weight_momentum: Vec<T>,
pub bias_momentum: T,
}
impl<T: Float> CandidateNeuron<T> {
pub fn new(
num_inputs: usize,
activation: ActivationFunction,
weight_range: (T, T),
random_seed: Option<u64>,
) -> Self {
let mut rng = if let Some(seed) = random_seed {
StdRng::seed_from_u64(seed)
} else {
StdRng::from_entropy()
};
use rand::Rng;
let (min_weight, max_weight) = weight_range;
let weights: Vec<T> = (0..num_inputs)
.map(|_| {
let w: f64 =
rng.gen_range(min_weight.to_f64().unwrap()..=max_weight.to_f64().unwrap());
T::from(w).unwrap()
})
.collect();
let bias_val: f64 =
rng.gen_range(min_weight.to_f64().unwrap()..=max_weight.to_f64().unwrap());
let bias = T::from(bias_val).unwrap();
Self {
weights,
bias,
activation,
steepness: T::one(),
correlation: T::zero(),
training_history: Vec::new(),
output: T::zero(),
weight_gradients: vec![T::zero(); num_inputs],
bias_gradient: T::zero(),
weight_momentum: vec![T::zero(); num_inputs],
bias_momentum: T::zero(),
}
}
pub fn calculate_output(&self, inputs: &[T]) -> T {
let mut sum = self.bias;
for (weight, input) in self.weights.iter().zip(inputs.iter()) {
sum = sum + *weight * *input;
}
match self.activation {
ActivationFunction::Sigmoid => {
let exp_neg = (-sum * self.steepness).exp();
T::one() / (T::one() + exp_neg)
}
ActivationFunction::Tanh => (sum * self.steepness).tanh(),
ActivationFunction::Linear => sum * self.steepness,
ActivationFunction::Gaussian => {
let neg_sq = -(sum * self.steepness) * (sum * self.steepness);
neg_sq.exp()
}
ActivationFunction::ReLU => {
if sum > T::zero() {
sum * self.steepness
} else {
T::zero()
}
}
_ => sum * self.steepness, }
}
pub fn activation_derivative(&self, output: T) -> T {
match self.activation {
ActivationFunction::Sigmoid => output * (T::one() - output) * self.steepness,
ActivationFunction::Tanh => (T::one() - output * output) * self.steepness,
ActivationFunction::Linear => self.steepness,
ActivationFunction::Gaussian => {
let neg_two = T::from(-2.0).unwrap();
neg_two * output * self.steepness * self.steepness
}
ActivationFunction::ReLU => {
if output > T::zero() {
self.steepness
} else {
T::zero()
}
}
_ => self.steepness, }
}
pub fn update_weights(&mut self, learning_rate: T, use_momentum: bool, momentum: T) {
for i in 0..self.weights.len() {
if use_momentum {
self.weight_momentum[i] =
momentum * self.weight_momentum[i] + learning_rate * self.weight_gradients[i];
self.weights[i] = self.weights[i] - self.weight_momentum[i];
} else {
self.weights[i] = self.weights[i] - learning_rate * self.weight_gradients[i];
}
}
if use_momentum {
self.bias_momentum = momentum * self.bias_momentum + learning_rate * self.bias_gradient;
self.bias = self.bias - self.bias_momentum;
} else {
self.bias = self.bias - learning_rate * self.bias_gradient;
}
}
}
pub struct CascadeTrainer<T: Float> {
pub config: CascadeConfig<T>,
pub network: Network<T>,
pub training_data: TrainingData<T>,
pub hidden_count: usize,
pub training_history: Vec<CascadeTrainingRecord<T>>,
pub current_epoch: usize,
pub best_error: T,
pub rng: rand::rngs::StdRng,
pub metrics: CascadeMetrics,
}
#[derive(Debug, Clone)]
pub struct CascadeTrainingRecord<T: Float> {
pub hidden_neuron_index: usize,
pub output_training_epochs: usize,
pub candidate_training_epochs: usize,
pub final_output_error: T,
pub best_candidate_correlation: T,
pub selected_activation: ActivationFunction,
pub convergence_reason: String,
}
#[derive(Debug, Clone)]
pub struct CascadeMetrics {
pub total_training_time: std::time::Duration,
pub output_training_time: std::time::Duration,
pub candidate_training_time: std::time::Duration,
pub correlation_calculation_time: std::time::Duration,
pub network_modification_time: std::time::Duration,
pub total_forward_passes: usize,
pub total_backward_passes: usize,
pub memory_usage_mb: f64,
pub peak_memory_usage_mb: f64,
}
impl Default for CascadeMetrics {
fn default() -> Self {
Self {
total_training_time: std::time::Duration::new(0, 0),
output_training_time: std::time::Duration::new(0, 0),
candidate_training_time: std::time::Duration::new(0, 0),
correlation_calculation_time: std::time::Duration::new(0, 0),
network_modification_time: std::time::Duration::new(0, 0),
total_forward_passes: 0,
total_backward_passes: 0,
memory_usage_mb: 0.0,
peak_memory_usage_mb: 0.0,
}
}
}
impl<T: Float> CascadeTrainer<T> {
pub fn new(
config: CascadeConfig<T>,
initial_network: Network<T>,
training_data: TrainingData<T>,
) -> Result<Self, CascadeError> {
Self::validate_config(&config)?;
Self::validate_training_data(&training_data, &initial_network)?;
let rng = if let Some(seed) = config.random_seed {
StdRng::seed_from_u64(seed)
} else {
StdRng::from_entropy()
};
Ok(Self {
config,
network: initial_network,
training_data,
hidden_count: 0,
training_history: Vec::new(),
current_epoch: 0,
best_error: T::infinity(),
rng,
metrics: CascadeMetrics::default(),
})
}
pub fn train(&mut self) -> Result<CascadeTrainingResult<T>, RuvFannError> {
let start_time = std::time::Instant::now();
#[cfg(feature = "logging")]
info!(
"Starting cascade correlation training with {} max hidden neurons",
self.config.max_hidden_neurons
);
self.train_output_weights()?;
while self.hidden_count < self.config.max_hidden_neurons {
if self.config.verbose {
println!(
"Adding hidden neuron {} of {}",
self.hidden_count + 1,
self.config.max_hidden_neurons
);
}
let best_candidate = self.train_candidates()?;
if best_candidate.correlation < self.config.min_correlation_improvement {
#[cfg(feature = "logging")]
info!("No candidate meets minimum correlation improvement threshold. Stopping cascade training.");
break;
}
self.install_candidate(best_candidate)?;
self.train_output_weights()?;
if self.best_error <= self.config.output_target_error {
#[cfg(feature = "logging")]
info!("Target error achieved. Stopping cascade training.");
break;
}
self.hidden_count += 1;
}
self.metrics.total_training_time = start_time.elapsed();
#[cfg(feature = "logging")]
info!(
"Cascade training completed. Added {} hidden neurons. Final error: {}",
self.hidden_count,
self.best_error.to_f64().unwrap_or(0.0)
);
Ok(CascadeTrainingResult {
final_network: self.network.clone(),
final_error: self.best_error,
hidden_neurons_added: self.hidden_count,
training_history: self.training_history.clone(),
metrics: self.metrics.clone(),
convergence_reason: self.determine_convergence_reason(),
})
}
fn train_output_weights(&mut self) -> Result<(), RuvFannError> {
let start_time = std::time::Instant::now();
#[cfg(feature = "logging")]
debug!(
"Training output weights for {} epochs",
self.config.output_max_epochs
);
let mut patience_counter = 0;
let mut best_epoch_error = T::infinity();
for epoch in 0..self.config.output_max_epochs {
let epoch_error = self.train_output_epoch()?;
if epoch_error < best_epoch_error {
best_epoch_error = epoch_error;
self.best_error = epoch_error;
patience_counter = 0;
} else {
patience_counter += 1;
}
if patience_counter >= self.config.patience {
#[cfg(feature = "logging")]
debug!("Early stopping output training at epoch {epoch} due to patience");
break;
}
if epoch_error <= self.config.output_target_error {
#[cfg(feature = "logging")]
debug!("Target error reached in output training at epoch {epoch}");
break;
}
if self.config.verbose && epoch % 100 == 0 {
println!(
"Output training epoch {}: error = {}",
epoch,
epoch_error.to_f64().unwrap_or(0.0)
);
}
}
self.metrics.output_training_time += start_time.elapsed();
Ok(())
}
fn train_output_epoch(&mut self) -> Result<T, RuvFannError> {
let mut total_error = T::zero();
let num_samples = T::from(self.training_data.inputs.len()).unwrap();
let inputs = self.training_data.inputs.clone();
let outputs = self.training_data.outputs.clone();
for (input, target) in inputs.iter().zip(outputs.iter()) {
let output = self.network.run(input);
self.metrics.total_forward_passes += 1;
let sample_error = self.calculate_output_error(&output, target);
total_error = total_error + sample_error;
self.update_output_weights(input, target, &output)?;
self.metrics.total_backward_passes += 1;
}
Ok(total_error / num_samples)
}
fn train_candidates(&mut self) -> Result<CandidateNeuron<T>, RuvFannError> {
let start_time = std::time::Instant::now();
#[cfg(feature = "logging")]
debug!("Training {} candidate neurons", self.config.num_candidates);
let mut candidates = self.generate_candidates()?;
#[cfg(feature = "parallel")]
{
if self.config.parallel_candidates {
self.train_candidates_sequential(&mut candidates)?;
} else {
self.train_candidates_sequential(&mut candidates)?;
}
}
#[cfg(not(feature = "parallel"))]
{
self.train_candidates_sequential(&mut candidates)?;
}
let best_candidate = candidates
.into_iter()
.max_by(|a, b| a.correlation.partial_cmp(&b.correlation).unwrap())
.ok_or_else(|| {
cascade_error!(
CascadeErrorCategory::CandidateSelection,
"No candidates generated"
)
})?;
self.metrics.candidate_training_time += start_time.elapsed();
#[cfg(feature = "logging")]
debug!(
"Best candidate correlation: {}",
best_candidate.correlation.to_f64().unwrap_or(0.0)
);
Ok(best_candidate)
}
fn generate_candidates(&mut self) -> Result<Vec<CandidateNeuron<T>>, RuvFannError> {
let num_inputs = self.calculate_candidate_input_size();
let mut candidates = Vec::with_capacity(self.config.num_candidates);
for _ in 0..self.config.num_candidates {
let activation_idx = self
.rng
.gen_range(0..self.config.candidate_activations.len());
let activation = self.config.candidate_activations[activation_idx];
let candidate = CandidateNeuron::new(
num_inputs,
activation,
self.config.candidate_weight_range,
self.config.random_seed,
);
candidates.push(candidate);
}
Ok(candidates)
}
fn train_candidates_sequential(
&mut self,
candidates: &mut [CandidateNeuron<T>],
) -> Result<(), RuvFannError> {
for candidate in candidates.iter_mut() {
self.train_single_candidate(candidate)?;
}
Ok(())
}
fn train_single_candidate(
&mut self,
candidate: &mut CandidateNeuron<T>,
) -> Result<(), RuvFannError> {
let mut best_correlation = T::zero();
let mut patience_counter = 0;
for _epoch in 0..self.config.candidate_max_epochs {
let residuals = self.calculate_residuals()?;
self.train_candidate_epoch(candidate, &residuals)?;
let correlation = self.calculate_correlation(candidate, &residuals)?;
candidate.correlation = correlation;
candidate.training_history.push(correlation);
if correlation > best_correlation {
best_correlation = correlation;
patience_counter = 0;
} else {
patience_counter += 1;
}
if patience_counter >= self.config.patience {
break;
}
if correlation >= self.config.candidate_target_correlation {
break;
}
}
Ok(())
}
fn calculate_residuals(&mut self) -> Result<Vec<Vec<T>>, RuvFannError> {
let mut residuals = Vec::with_capacity(self.training_data.inputs.len());
for (input, target) in self
.training_data
.inputs
.iter()
.zip(self.training_data.outputs.iter())
{
let output = self.network.run(input);
let residual: Vec<T> = output
.iter()
.zip(target.iter())
.map(|(&o, &t)| t - o)
.collect();
residuals.push(residual);
}
Ok(residuals)
}
fn calculate_correlation(
&mut self,
candidate: &mut CandidateNeuron<T>,
residuals: &[Vec<T>],
) -> Result<T, RuvFannError> {
let start_time = std::time::Instant::now();
let mut candidate_outputs = Vec::with_capacity(self.training_data.inputs.len());
for input in &self.training_data.inputs {
let candidate_input = self.extract_candidate_input(input);
let output = candidate.calculate_output(&candidate_input);
candidate_outputs.push(output);
}
let mut total_correlation = T::zero();
let num_outputs = self.training_data.outputs[0].len();
for output_idx in 0..num_outputs {
let residual_values: Vec<T> = residuals.iter().map(|r| r[output_idx]).collect();
let correlation = self.pearson_correlation(&candidate_outputs, &residual_values)?;
total_correlation = total_correlation + correlation.abs();
}
self.metrics.correlation_calculation_time += start_time.elapsed();
Ok(total_correlation)
}
fn pearson_correlation(&self, x: &[T], y: &[T]) -> Result<T, RuvFannError> {
if x.len() != y.len() || x.is_empty() {
return Err(cascade_error!(
CascadeErrorCategory::CorrelationCalculation,
"Invalid input arrays for correlation calculation"
));
}
let n = T::from(x.len()).unwrap();
let mean_x = x.iter().fold(T::zero(), |acc, &val| acc + val) / n;
let mean_y = y.iter().fold(T::zero(), |acc, &val| acc + val) / n;
let mut numerator = T::zero();
let mut sum_sq_x = T::zero();
let mut sum_sq_y = T::zero();
for (&xi, &yi) in x.iter().zip(y.iter()) {
let diff_x = xi - mean_x;
let diff_y = yi - mean_y;
numerator = numerator + diff_x * diff_y;
sum_sq_x = sum_sq_x + diff_x * diff_x;
sum_sq_y = sum_sq_y + diff_y * diff_y;
}
let denominator = (sum_sq_x * sum_sq_y).sqrt();
if denominator == T::zero() {
Ok(T::zero())
} else {
Ok(numerator / denominator)
}
}
fn install_candidate(&mut self, candidate: CandidateNeuron<T>) -> Result<(), RuvFannError> {
let start_time = std::time::Instant::now();
#[cfg(feature = "logging")]
debug!(
"Installing candidate neuron with correlation {}",
candidate.correlation.to_f64().unwrap_or(0.0)
);
let record = CascadeTrainingRecord {
hidden_neuron_index: self.hidden_count,
output_training_epochs: 0, candidate_training_epochs: candidate.training_history.len(),
final_output_error: self.best_error,
best_candidate_correlation: candidate.correlation,
selected_activation: candidate.activation,
convergence_reason: "Candidate installed".to_string(),
};
self.training_history.push(record);
self.metrics.network_modification_time += start_time.elapsed();
Ok(())
}
fn calculate_candidate_input_size(&self) -> usize {
self.network.num_inputs() + self.hidden_count
}
fn extract_candidate_input(&self, input: &[T]) -> Vec<T> {
input.to_vec()
}
fn calculate_output_error(&self, output: &[T], target: &[T]) -> T {
output
.iter()
.zip(target.iter())
.map(|(&o, &t)| (o - t) * (o - t))
.fold(T::zero(), |acc, err| acc + err)
}
fn update_output_weights(
&mut self,
_input: &[T],
_target: &[T],
_output: &[T],
) -> Result<(), RuvFannError> {
Ok(())
}
fn train_candidate_epoch(
&mut self,
_candidate: &mut CandidateNeuron<T>,
_residuals: &[Vec<T>],
) -> Result<(), RuvFannError> {
Ok(())
}
fn determine_convergence_reason(&self) -> String {
if self.best_error <= self.config.output_target_error {
"Target error achieved".to_string()
} else if self.hidden_count >= self.config.max_hidden_neurons {
"Maximum hidden neurons reached".to_string()
} else {
"No further improvement possible".to_string()
}
}
fn validate_config(config: &CascadeConfig<T>) -> Result<(), CascadeError> {
if config.max_hidden_neurons == 0 {
return Err(CascadeError::InvalidConfiguration(
"max_hidden_neurons must be greater than 0".to_string(),
));
}
if config.num_candidates == 0 {
return Err(CascadeError::InvalidConfiguration(
"num_candidates must be greater than 0".to_string(),
));
}
if config.output_learning_rate <= T::zero() {
return Err(CascadeError::InvalidConfiguration(
"output_learning_rate must be positive".to_string(),
));
}
if config.candidate_learning_rate <= T::zero() {
return Err(CascadeError::InvalidConfiguration(
"candidate_learning_rate must be positive".to_string(),
));
}
if config.candidate_activations.is_empty() {
return Err(CascadeError::InvalidConfiguration(
"candidate_activations cannot be empty".to_string(),
));
}
Ok(())
}
fn validate_training_data(
data: &TrainingData<T>,
network: &Network<T>,
) -> Result<(), CascadeError> {
if data.inputs.is_empty() {
return Err(CascadeError::InvalidConfiguration(
"Training data cannot be empty".to_string(),
));
}
if data.inputs.len() != data.outputs.len() {
return Err(CascadeError::InvalidConfiguration(
"Input and output data must have same length".to_string(),
));
}
let expected_inputs = network.num_inputs();
let expected_outputs = network.num_outputs();
for (i, input) in data.inputs.iter().enumerate() {
if input.len() != expected_inputs {
return Err(CascadeError::InvalidConfiguration(format!(
"Input {} has wrong size: expected {}, got {}",
i,
expected_inputs,
input.len()
)));
}
}
for (i, output) in data.outputs.iter().enumerate() {
if output.len() != expected_outputs {
return Err(CascadeError::InvalidConfiguration(format!(
"Output {} has wrong size: expected {}, got {}",
i,
expected_outputs,
output.len()
)));
}
}
Ok(())
}
#[cfg(feature = "parallel")]
#[allow(dead_code)]
fn train_single_candidate_parallel(
&self,
_candidate: &mut CandidateNeuron<T>,
_training_data: &TrainingData<T>,
_config: &CascadeConfig<T>,
) -> Result<(), RuvFannError> {
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct CascadeTrainingResult<T: Float> {
pub final_network: Network<T>,
pub final_error: T,
pub hidden_neurons_added: usize,
pub training_history: Vec<CascadeTrainingRecord<T>>,
pub metrics: CascadeMetrics,
pub convergence_reason: String,
}
pub struct CascadeBuilder<T: Float> {
config: CascadeConfig<T>,
}
impl<T: Float> CascadeBuilder<T> {
pub fn new() -> Self {
Self {
config: CascadeConfig::default(),
}
}
pub fn max_hidden_neurons(mut self, max: usize) -> Self {
self.config.max_hidden_neurons = max;
self
}
pub fn num_candidates(mut self, num: usize) -> Self {
self.config.num_candidates = num;
self
}
pub fn output_learning_rate(mut self, rate: T) -> Self {
self.config.output_learning_rate = rate;
self
}
pub fn candidate_learning_rate(mut self, rate: T) -> Self {
self.config.candidate_learning_rate = rate;
self
}
pub fn target_error(mut self, error: T) -> Self {
self.config.output_target_error = error;
self
}
pub fn parallel_candidates(mut self, enabled: bool) -> Self {
self.config.parallel_candidates = enabled;
self
}
pub fn verbose(mut self, enabled: bool) -> Self {
self.config.verbose = enabled;
self
}
pub fn random_seed(mut self, seed: u64) -> Self {
self.config.random_seed = Some(seed);
self
}
pub fn build(self) -> CascadeConfig<T> {
self.config
}
}
impl<T: Float> Default for CascadeBuilder<T> {
fn default() -> Self {
Self::new()
}
}
#[cfg(feature = "parallel")]
impl<T: Float + Send + Sync> CascadeTrainer<T>
where
T::FromStrRadixErr: Send + Sync,
{
#[allow(dead_code)]
fn train_candidates_parallel(
&mut self,
candidates: &mut [CandidateNeuron<T>],
) -> Result<(), RuvFannError> {
use rayon::prelude::*;
use std::sync::Arc;
let training_data = Arc::new(self.training_data.clone());
let config = Arc::new(self.config.clone());
let network = Arc::new(self.network.clone());
let results: Vec<Result<(), RuvFannError>> = candidates
.par_iter_mut()
.map(|candidate| {
let local_data = training_data.clone();
let local_config = config.clone();
let local_network = network.clone();
self.train_single_candidate_with_data(
candidate,
&*local_data,
&*local_config,
&*local_network,
)
})
.collect();
for result in results {
result?;
}
Ok(())
}
#[allow(dead_code)]
fn train_single_candidate_with_data(
&self,
candidate: &mut CandidateNeuron<T>,
data: &TrainingData<T>,
config: &CascadeConfig<T>,
network: &Network<T>,
) -> Result<(), RuvFannError> {
let mut best_correlation = T::zero();
for _epoch in 0..config.candidate_max_epochs {
let correlation =
self.calculate_candidate_correlation_with_data(candidate, data, network)?;
if correlation > best_correlation {
best_correlation = correlation;
candidate.correlation = correlation;
}
if correlation >= config.candidate_target_correlation {
break;
}
self.update_candidate_weights_simple(
candidate,
data,
network,
config.candidate_learning_rate,
)?;
}
Ok(())
}
#[allow(dead_code)]
fn calculate_candidate_correlation_with_data(
&self,
candidate: &CandidateNeuron<T>,
data: &TrainingData<T>,
_network: &Network<T>,
) -> Result<T, RuvFannError> {
let mut sum = T::zero();
let mut count = 0;
for (input, _target) in data.inputs.iter().zip(data.outputs.iter()) {
let candidate_output = candidate.calculate_output(input);
sum = sum + candidate_output.abs();
count += 1;
}
if count == 0 {
Ok(T::zero())
} else {
Ok(sum / T::from(count).unwrap())
}
}
#[allow(dead_code)]
fn update_candidate_weights_simple(
&self,
candidate: &mut CandidateNeuron<T>,
_data: &TrainingData<T>,
_network: &Network<T>,
learning_rate: T,
) -> Result<(), RuvFannError> {
for i in 0..candidate.weights.len() {
let delta = T::from(0.01).unwrap() * learning_rate;
candidate.weights[i] = candidate.weights[i] - delta;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::NetworkBuilder;
#[test]
fn test_cascade_config_default() {
let config: CascadeConfig<f32> = CascadeConfig::default();
assert_eq!(config.max_hidden_neurons, 100);
assert_eq!(config.num_candidates, 8);
assert!(config.output_learning_rate > 0.0);
}
#[test]
fn test_cascade_builder() {
let config: CascadeConfig<f32> = CascadeBuilder::new()
.max_hidden_neurons(50)
.num_candidates(4)
.target_error(0.001)
.verbose(true)
.build();
assert_eq!(config.max_hidden_neurons, 50);
assert_eq!(config.num_candidates, 4);
assert_eq!(config.output_target_error, 0.001);
assert!(config.verbose);
}
#[test]
fn test_candidate_neuron_creation() {
let candidate: CandidateNeuron<f32> =
CandidateNeuron::new(5, ActivationFunction::Sigmoid, (-1.0, 1.0), Some(42));
assert_eq!(candidate.weights.len(), 5);
assert_eq!(candidate.activation, ActivationFunction::Sigmoid);
assert_eq!(candidate.correlation, 0.0);
}
#[test]
fn test_config_validation() {
let config: CascadeConfig<f32> = CascadeConfig {
max_hidden_neurons: 0,
..CascadeConfig::default()
};
assert!(CascadeTrainer::validate_config(&config).is_err());
}
#[test]
fn test_pearson_correlation() {
let network = NetworkBuilder::<f32>::new()
.input_layer(2)
.output_layer(1)
.build();
let training_data = TrainingData {
inputs: vec![vec![0.0, 0.0], vec![1.0, 1.0]],
outputs: vec![vec![0.0], vec![1.0]],
};
let config = CascadeConfig::default();
let trainer = CascadeTrainer::new(config, network, training_data).unwrap();
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let y = vec![2.0, 4.0, 6.0, 8.0, 10.0];
let correlation = trainer.pearson_correlation(&x, &y).unwrap();
assert!((correlation - 1.0).abs() < 1e-6); }
}