1#![allow(clippy::needless_range_loop)]
4
5use super::*;
6use num_traits::Float;
7use std::collections::HashMap;
8
9pub 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 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 self.weight_step_sizes = network
70 .layers
71 .iter()
72 .skip(1) .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) .map(|layer| vec![self.delta_zero; layer.neurons.len()])
89 .collect();
90
91 self.previous_weight_gradients = network
93 .layers
94 .iter()
95 .skip(1) .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) .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 (step_size * self.increase_factor).min(self.delta_max)
123 } else if sign_change < T::zero() {
124 (step_size * self.decrease_factor).max(self.delta_min)
126 } else {
127 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 let simple_network = network_to_simple(network);
153
154 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 for (input, desired_output) in data.inputs.iter().zip(data.outputs.iter()) {
168 let activations = forward_propagate(&simple_network, input);
170
171 let output = &activations[activations.len() - 1];
173
174 total_error = total_error + self.error_function.calculate(output, desired_output);
176
177 let (weight_gradients, bias_gradients) = calculate_gradients(
179 &simple_network,
180 &activations,
181 desired_output,
182 self.error_function.as_ref(),
183 );
184
185 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 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 let mut weight_updates = Vec::new();
213 let mut bias_updates = Vec::new();
214
215 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 if sign_change > T::zero() {
226 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 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 self.previous_weight_gradients[layer_idx][i] = current_gradient;
243 } else if sign_change < T::zero() {
244 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 layer_weight_updates.push(T::zero());
251
252 self.previous_weight_gradients[layer_idx][i] = T::zero();
254 } else {
255 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 self.previous_weight_gradients[layer_idx][i] = current_gradient;
267 }
268 }
269
270 weight_updates.push(layer_weight_updates);
271 }
272
273 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 if sign_change > T::zero() {
284 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 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 self.previous_bias_gradients[layer_idx][i] = current_gradient;
301 } else if sign_change < T::zero() {
302 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 layer_bias_updates.push(T::zero());
309
310 self.previous_bias_gradients[layer_idx][i] = T::zero();
312 } else {
313 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 self.previous_bias_gradients[layer_idx][i] = current_gradient;
325 }
326 }
327
328 bias_updates.push(layer_bias_updates);
329 }
330
331 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 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 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 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 }
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}