Skip to main content

rfann/training/
rprop.rs

1//! Resilient Propagation (RPROP) training algorithm
2
3#![allow(clippy::needless_range_loop)]
4
5use super::*;
6use num_traits::Float;
7use std::collections::HashMap;
8
9/// RPROP (Resilient Propagation) trainer
10/// An adaptive learning algorithm that only uses the sign of the gradient
11pub struct Rprop<T: Float + Send + Default> {
12    increase_factor: T,
13    decrease_factor: T,
14    delta_min: T,
15    delta_max: T,
16    delta_zero: T,
17    error_function: Box<dyn ErrorFunction<T>>,
18
19    // State variables
20    weight_step_sizes: Vec<Vec<T>>,
21    bias_step_sizes: Vec<Vec<T>>,
22    previous_weight_gradients: Vec<Vec<T>>,
23    previous_bias_gradients: Vec<Vec<T>>,
24
25    callback: Option<TrainingCallback<T>>,
26}
27
28impl<T: Float + Send + Default> Rprop<T> {
29    pub fn new() -> Self {
30        Self {
31            increase_factor: T::from(1.2).unwrap(),
32            decrease_factor: T::from(0.5).unwrap(),
33            delta_min: T::zero(),
34            delta_max: T::from(50.0).unwrap(),
35            delta_zero: T::from(0.1).unwrap(),
36            error_function: Box::new(MseError),
37            weight_step_sizes: Vec::new(),
38            bias_step_sizes: Vec::new(),
39            previous_weight_gradients: Vec::new(),
40            previous_bias_gradients: Vec::new(),
41            callback: None,
42        }
43    }
44
45    pub fn with_parameters(
46        mut self,
47        increase_factor: T,
48        decrease_factor: T,
49        delta_min: T,
50        delta_max: T,
51        delta_zero: T,
52    ) -> Self {
53        self.increase_factor = increase_factor;
54        self.decrease_factor = decrease_factor;
55        self.delta_min = delta_min;
56        self.delta_max = delta_max;
57        self.delta_zero = delta_zero;
58        self
59    }
60
61    pub fn with_error_function(mut self, error_function: Box<dyn ErrorFunction<T>>) -> Self {
62        self.error_function = error_function;
63        self
64    }
65
66    fn initialize_state(&mut self, network: &Network<T>) {
67        if self.weight_step_sizes.is_empty() {
68            // Initialize step sizes and gradients for each layer
69            self.weight_step_sizes = network
70                .layers
71                .iter()
72                .skip(1) // Skip input layer
73                .map(|layer| {
74                    let num_neurons = layer.neurons.len();
75                    let num_connections = if layer.neurons.is_empty() {
76                        0
77                    } else {
78                        layer.neurons[0].connections.len()
79                    };
80                    vec![self.delta_zero; num_neurons * num_connections]
81                })
82                .collect();
83
84            self.bias_step_sizes = network
85                .layers
86                .iter()
87                .skip(1) // Skip input layer
88                .map(|layer| vec![self.delta_zero; layer.neurons.len()])
89                .collect();
90
91            // Initialize previous gradients to zero
92            self.previous_weight_gradients = network
93                .layers
94                .iter()
95                .skip(1) // Skip input layer
96                .map(|layer| {
97                    let num_neurons = layer.neurons.len();
98                    let num_connections = if layer.neurons.is_empty() {
99                        0
100                    } else {
101                        layer.neurons[0].connections.len()
102                    };
103                    vec![T::zero(); num_neurons * num_connections]
104                })
105                .collect();
106
107            self.previous_bias_gradients = network
108                .layers
109                .iter()
110                .skip(1) // Skip input layer
111                .map(|layer| vec![T::zero(); layer.neurons.len()])
112                .collect();
113        }
114    }
115
116    #[allow(dead_code)]
117    fn update_step_size(&self, step_size: T, gradient: T, previous_gradient: T) -> T {
118        let sign_change = gradient * previous_gradient;
119
120        if sign_change > T::zero() {
121            // Same sign: increase step size
122            (step_size * self.increase_factor).min(self.delta_max)
123        } else if sign_change < T::zero() {
124            // Different sign: decrease step size
125            (step_size * self.decrease_factor).max(self.delta_min)
126        } else {
127            // No previous gradient or zero gradient: keep step size
128            step_size
129        }
130    }
131}
132
133impl<T: Float + Send + Default> Default for Rprop<T> {
134    fn default() -> Self {
135        Self::new()
136    }
137}
138
139impl<T: Float + Send + Default> TrainingAlgorithm<T> for Rprop<T> {
140    fn train_epoch(
141        &mut self,
142        network: &mut Network<T>,
143        data: &TrainingData<T>,
144    ) -> Result<T, TrainingError> {
145        use super::helpers::*;
146
147        self.initialize_state(network);
148
149        let mut total_error = T::zero();
150
151        // Convert network to simplified form for easier manipulation
152        let simple_network = network_to_simple(network);
153
154        // Initialize gradient accumulators
155        let mut accumulated_weight_gradients = simple_network
156            .weights
157            .iter()
158            .map(|w| vec![T::zero(); w.len()])
159            .collect::<Vec<_>>();
160        let mut accumulated_bias_gradients = simple_network
161            .biases
162            .iter()
163            .map(|b| vec![T::zero(); b.len()])
164            .collect::<Vec<_>>();
165
166        // Calculate gradients over entire dataset
167        for (input, desired_output) in data.inputs.iter().zip(data.outputs.iter()) {
168            // Forward propagation to get all layer activations
169            let activations = forward_propagate(&simple_network, input);
170
171            // Get output from last layer
172            let output = &activations[activations.len() - 1];
173
174            // Calculate error
175            total_error = total_error + self.error_function.calculate(output, desired_output);
176
177            // Calculate gradients using backpropagation
178            let (weight_gradients, bias_gradients) = calculate_gradients(
179                &simple_network,
180                &activations,
181                desired_output,
182                self.error_function.as_ref(),
183            );
184
185            // Accumulate gradients
186            for layer_idx in 0..weight_gradients.len() {
187                for i in 0..weight_gradients[layer_idx].len() {
188                    accumulated_weight_gradients[layer_idx][i] =
189                        accumulated_weight_gradients[layer_idx][i] + weight_gradients[layer_idx][i];
190                }
191                for i in 0..bias_gradients[layer_idx].len() {
192                    accumulated_bias_gradients[layer_idx][i] =
193                        accumulated_bias_gradients[layer_idx][i] + bias_gradients[layer_idx][i];
194                }
195            }
196        }
197
198        // Average gradients by batch size
199        let batch_size = T::from(data.inputs.len()).unwrap();
200        for layer_idx in 0..accumulated_weight_gradients.len() {
201            for i in 0..accumulated_weight_gradients[layer_idx].len() {
202                accumulated_weight_gradients[layer_idx][i] =
203                    accumulated_weight_gradients[layer_idx][i] / batch_size;
204            }
205            for i in 0..accumulated_bias_gradients[layer_idx].len() {
206                accumulated_bias_gradients[layer_idx][i] =
207                    accumulated_bias_gradients[layer_idx][i] / batch_size;
208            }
209        }
210
211        // Apply RPROP updates
212        let mut weight_updates = Vec::new();
213        let mut bias_updates = Vec::new();
214
215        // Update weights using RPROP algorithm
216        for layer_idx in 0..accumulated_weight_gradients.len() {
217            let mut layer_weight_updates = Vec::new();
218
219            for i in 0..accumulated_weight_gradients[layer_idx].len() {
220                let current_gradient = accumulated_weight_gradients[layer_idx][i];
221                let previous_gradient = self.previous_weight_gradients[layer_idx][i];
222                let sign_change = current_gradient * previous_gradient;
223
224                // Update step size based on gradient sign change
225                if sign_change > T::zero() {
226                    // Same sign - increase step size
227                    self.weight_step_sizes[layer_idx][i] = (self.weight_step_sizes[layer_idx][i]
228                        * self.increase_factor)
229                        .min(self.delta_max);
230
231                    // Update weight
232                    let update = if current_gradient > T::zero() {
233                        -self.weight_step_sizes[layer_idx][i]
234                    } else if current_gradient < T::zero() {
235                        self.weight_step_sizes[layer_idx][i]
236                    } else {
237                        T::zero()
238                    };
239                    layer_weight_updates.push(update);
240
241                    // Store gradient for next iteration
242                    self.previous_weight_gradients[layer_idx][i] = current_gradient;
243                } else if sign_change < T::zero() {
244                    // Sign changed - decrease step size and backtrack
245                    self.weight_step_sizes[layer_idx][i] = (self.weight_step_sizes[layer_idx][i]
246                        * self.decrease_factor)
247                        .max(self.delta_min);
248
249                    // Don't update weight (backtrack)
250                    layer_weight_updates.push(T::zero());
251
252                    // Set gradient to zero to prevent another sign change detection
253                    self.previous_weight_gradients[layer_idx][i] = T::zero();
254                } else {
255                    // No previous gradient or current gradient is zero
256                    let update = if current_gradient > T::zero() {
257                        -self.weight_step_sizes[layer_idx][i]
258                    } else if current_gradient < T::zero() {
259                        self.weight_step_sizes[layer_idx][i]
260                    } else {
261                        T::zero()
262                    };
263                    layer_weight_updates.push(update);
264
265                    // Store gradient for next iteration
266                    self.previous_weight_gradients[layer_idx][i] = current_gradient;
267                }
268            }
269
270            weight_updates.push(layer_weight_updates);
271        }
272
273        // Update biases using RPROP algorithm
274        for layer_idx in 0..accumulated_bias_gradients.len() {
275            let mut layer_bias_updates = Vec::new();
276
277            for i in 0..accumulated_bias_gradients[layer_idx].len() {
278                let current_gradient = accumulated_bias_gradients[layer_idx][i];
279                let previous_gradient = self.previous_bias_gradients[layer_idx][i];
280                let sign_change = current_gradient * previous_gradient;
281
282                // Update step size based on gradient sign change
283                if sign_change > T::zero() {
284                    // Same sign - increase step size
285                    self.bias_step_sizes[layer_idx][i] = (self.bias_step_sizes[layer_idx][i]
286                        * self.increase_factor)
287                        .min(self.delta_max);
288
289                    // Update bias
290                    let update = if current_gradient > T::zero() {
291                        -self.bias_step_sizes[layer_idx][i]
292                    } else if current_gradient < T::zero() {
293                        self.bias_step_sizes[layer_idx][i]
294                    } else {
295                        T::zero()
296                    };
297                    layer_bias_updates.push(update);
298
299                    // Store gradient for next iteration
300                    self.previous_bias_gradients[layer_idx][i] = current_gradient;
301                } else if sign_change < T::zero() {
302                    // Sign changed - decrease step size and backtrack
303                    self.bias_step_sizes[layer_idx][i] = (self.bias_step_sizes[layer_idx][i]
304                        * self.decrease_factor)
305                        .max(self.delta_min);
306
307                    // Don't update bias (backtrack)
308                    layer_bias_updates.push(T::zero());
309
310                    // Set gradient to zero to prevent another sign change detection
311                    self.previous_bias_gradients[layer_idx][i] = T::zero();
312                } else {
313                    // No previous gradient or current gradient is zero
314                    let update = if current_gradient > T::zero() {
315                        -self.bias_step_sizes[layer_idx][i]
316                    } else if current_gradient < T::zero() {
317                        self.bias_step_sizes[layer_idx][i]
318                    } else {
319                        T::zero()
320                    };
321                    layer_bias_updates.push(update);
322
323                    // Store gradient for next iteration
324                    self.previous_bias_gradients[layer_idx][i] = current_gradient;
325                }
326            }
327
328            bias_updates.push(layer_bias_updates);
329        }
330
331        // Apply the updates to the actual network
332        apply_updates_to_network(network, &weight_updates, &bias_updates);
333
334        Ok(total_error / batch_size)
335    }
336
337    fn calculate_error(&self, network: &Network<T>, data: &TrainingData<T>) -> T {
338        let mut total_error = T::zero();
339        let mut network_clone = network.clone();
340
341        for (input, desired_output) in data.inputs.iter().zip(data.outputs.iter()) {
342            let output = network_clone.run(input);
343            total_error = total_error + self.error_function.calculate(&output, desired_output);
344        }
345
346        total_error / T::from(data.inputs.len()).unwrap()
347    }
348
349    fn count_bit_fails(
350        &self,
351        network: &Network<T>,
352        data: &TrainingData<T>,
353        bit_fail_limit: T,
354    ) -> usize {
355        let mut bit_fails = 0;
356        let mut network_clone = network.clone();
357
358        for (input, desired_output) in data.inputs.iter().zip(data.outputs.iter()) {
359            let output = network_clone.run(input);
360
361            for (&actual, &desired) in output.iter().zip(desired_output.iter()) {
362                if (actual - desired).abs() > bit_fail_limit {
363                    bit_fails += 1;
364                }
365            }
366        }
367
368        bit_fails
369    }
370
371    fn save_state(&self) -> TrainingState<T> {
372        let mut state = HashMap::new();
373
374        // Save RPROP parameters
375        state.insert("increase_factor".to_string(), vec![self.increase_factor]);
376        state.insert("decrease_factor".to_string(), vec![self.decrease_factor]);
377        state.insert("delta_min".to_string(), vec![self.delta_min]);
378        state.insert("delta_max".to_string(), vec![self.delta_max]);
379        state.insert("delta_zero".to_string(), vec![self.delta_zero]);
380
381        // Save step sizes (flattened)
382        let mut all_weight_steps = Vec::new();
383        for layer_steps in &self.weight_step_sizes {
384            all_weight_steps.extend_from_slice(layer_steps);
385        }
386        state.insert("weight_step_sizes".to_string(), all_weight_steps);
387
388        let mut all_bias_steps = Vec::new();
389        for layer_steps in &self.bias_step_sizes {
390            all_bias_steps.extend_from_slice(layer_steps);
391        }
392        state.insert("bias_step_sizes".to_string(), all_bias_steps);
393
394        TrainingState {
395            epoch: 0,
396            best_error: T::from(f32::MAX).unwrap(),
397            algorithm_specific: state,
398        }
399    }
400
401    fn restore_state(&mut self, state: TrainingState<T>) {
402        // Restore RPROP parameters
403        if let Some(val) = state.algorithm_specific.get("increase_factor") {
404            if !val.is_empty() {
405                self.increase_factor = val[0];
406            }
407        }
408        if let Some(val) = state.algorithm_specific.get("decrease_factor") {
409            if !val.is_empty() {
410                self.decrease_factor = val[0];
411            }
412        }
413        if let Some(val) = state.algorithm_specific.get("delta_min") {
414            if !val.is_empty() {
415                self.delta_min = val[0];
416            }
417        }
418        if let Some(val) = state.algorithm_specific.get("delta_max") {
419            if !val.is_empty() {
420                self.delta_max = val[0];
421            }
422        }
423        if let Some(val) = state.algorithm_specific.get("delta_zero") {
424            if !val.is_empty() {
425                self.delta_zero = val[0];
426            }
427        }
428
429        // Note: Step sizes would need network structure info to properly restore
430        // This is a simplified version - in production, you'd need to store layer sizes too
431    }
432
433    fn set_callback(&mut self, callback: TrainingCallback<T>) {
434        self.callback = Some(callback);
435    }
436
437    fn call_callback(
438        &mut self,
439        epoch: usize,
440        network: &Network<T>,
441        data: &TrainingData<T>,
442    ) -> bool {
443        let error = self.calculate_error(network, data);
444        if let Some(ref mut callback) = self.callback {
445            callback(epoch, error)
446        } else {
447            true
448        }
449    }
450}