1use crate::error::{OptimError, Result};
7use crate::utils::{scalar_or, try_scalar};
8use scirs2_core::ndarray::{Array, Dimension, ScalarOperand, Zip};
9use scirs2_core::numeric::Float;
10use std::collections::VecDeque;
11use std::fmt::Debug;
12
13#[derive(Debug, Clone, Copy, PartialEq)]
15pub enum AveragingMethod {
16 MovingAverage,
18 ExponentialMovingAverage {
20 decay: f64,
22 },
23 StochasticWeightAveraging,
25 ModelSoup,
27}
28
29#[derive(Debug)]
31pub struct WeightAverager<A: Float, D: Dimension> {
32 averaged_weights: Vec<Array<A, D>>,
34 weight_history: VecDeque<Vec<Array<A, D>>>,
36 step_count: usize,
38 method: AveragingMethod,
40 max_history: usize,
42 initialized: bool,
44 ema_decay: A,
46}
47
48impl<A: Float + ScalarOperand + Debug, D: Dimension + Send + Sync> WeightAverager<A, D> {
49 pub fn new(method: AveragingMethod, maxhistory: usize) -> Self {
51 let ema_decay = match method {
52 AveragingMethod::ExponentialMovingAverage { decay } => {
53 A::from(decay).unwrap_or_else(|| scalar_or(0.999, A::zero()))
54 }
55 _ => scalar_or(0.999, A::zero()),
56 };
57
58 Self {
59 averaged_weights: Vec::new(),
60 weight_history: VecDeque::new(),
61 step_count: 0,
62 method,
63 max_history: maxhistory,
64 initialized: false,
65 ema_decay,
66 }
67 }
68
69 pub fn initialize(&mut self, weights: &[Array<A, D>]) -> Result<()> {
71 if self.initialized {
72 return Err(OptimError::InvalidConfig(
73 "Weight averager already initialized".to_string(),
74 ));
75 }
76
77 self.averaged_weights = weights.to_vec();
78 self.initialized = true;
79 Ok(())
80 }
81
82 pub fn update(&mut self, weights: &[Array<A, D>]) -> Result<()> {
84 if !self.initialized {
85 self.initialize(weights)?;
86 return Ok(());
87 }
88
89 if weights.len() != self.averaged_weights.len() {
90 return Err(OptimError::DimensionMismatch(format!(
91 "Expected {} weight arrays, got {}",
92 self.averaged_weights.len(),
93 weights.len()
94 )));
95 }
96
97 self.step_count += 1;
98
99 match self.method {
100 AveragingMethod::MovingAverage => {
101 self.update_moving_average(weights)?;
102 }
103 AveragingMethod::ExponentialMovingAverage { .. } => {
104 self.update_exponential_moving_average(weights)?;
105 }
106 AveragingMethod::StochasticWeightAveraging => {
107 self.update_swa(weights)?;
108 }
109 AveragingMethod::ModelSoup => {
110 self.update_model_soup(weights)?;
111 }
112 }
113
114 Ok(())
115 }
116
117 fn update_moving_average(&mut self, weights: &[Array<A, D>]) -> Result<()> {
119 self.weight_history.push_back(weights.to_vec());
121
122 if self.weight_history.len() > self.max_history {
124 self.weight_history.pop_front();
125 }
126
127 self.compute_moving_average()
129 }
130
131 fn compute_moving_average(&mut self) -> Result<()> {
133 if self.weight_history.is_empty() {
134 return Ok(());
135 }
136
137 let num_snapshots = self.weight_history.len();
138 let inv_count = A::one() / try_scalar::<A, _>(num_snapshots)?;
139
140 for avg_weight in &mut self.averaged_weights {
142 avg_weight.fill(A::zero());
143 }
144
145 for snapshot in &self.weight_history {
147 for (avg_weight, weight) in self.averaged_weights.iter_mut().zip(snapshot.iter()) {
148 Zip::from(avg_weight).and(weight).for_each(|avg, &w| {
149 *avg = *avg + w;
150 });
151 }
152 }
153
154 for avg_weight in &mut self.averaged_weights {
156 avg_weight.mapv_inplace(|x| x * inv_count);
157 }
158
159 Ok(())
160 }
161
162 fn update_exponential_moving_average(&mut self, weights: &[Array<A, D>]) -> Result<()> {
164 let alpha = A::one() - self.ema_decay;
165
166 for (avg_weight, weight) in self.averaged_weights.iter_mut().zip(weights.iter()) {
167 Zip::from(avg_weight).and(weight).for_each(|avg, &w| {
168 *avg = self.ema_decay * *avg + alpha * w;
169 });
170 }
171
172 Ok(())
173 }
174
175 fn update_swa(&mut self, weights: &[Array<A, D>]) -> Result<()> {
177 let n = try_scalar::<A, _>(self.step_count)?;
179 let inv_n = A::one() / n;
180 let prev_weight = (n - A::one()) / n;
181
182 for (avg_weight, weight) in self.averaged_weights.iter_mut().zip(weights.iter()) {
183 Zip::from(avg_weight).and(weight).for_each(|avg, &w| {
184 *avg = prev_weight * *avg + inv_n * w;
185 });
186 }
187
188 Ok(())
189 }
190
191 fn update_model_soup(&mut self, weights: &[Array<A, D>]) -> Result<()> {
193 self.weight_history.push_back(weights.to_vec());
195
196 if self.weight_history.len() > self.max_history {
197 self.weight_history.pop_front();
198 }
199
200 self.compute_moving_average()
202 }
203
204 pub fn get_averaged_weights(&self) -> &[Array<A, D>] {
206 &self.averaged_weights
207 }
208
209 pub fn get_averaged_weights_cloned(&self) -> Vec<Array<A, D>> {
211 self.averaged_weights.clone()
212 }
213
214 pub fn reset(&mut self) {
216 self.weight_history.clear();
217 self.step_count = 0;
218 for weight in &mut self.averaged_weights {
219 weight.fill(A::zero());
220 }
221 }
222
223 pub fn step_count(&self) -> usize {
225 self.step_count
226 }
227
228 pub fn is_initialized(&self) -> bool {
230 self.initialized
231 }
232
233 pub fn method(&self) -> AveragingMethod {
235 self.method
236 }
237
238 pub fn set_ema_decay(&mut self, decay: A) {
240 self.ema_decay = decay;
241 }
242}
243
244#[derive(Debug)]
246pub struct PolyakAverager<A: Float, D: Dimension> {
247 averager: WeightAverager<A, D>,
249 initial_decay: A,
251 final_decay: A,
253 decay_steps: usize,
255}
256
257impl<A: Float + ScalarOperand + Debug, D: Dimension + Send + Sync> PolyakAverager<A, D> {
258 pub fn new(initial_decay: A, final_decay: A, decaysteps: usize) -> Self {
260 let method = AveragingMethod::ExponentialMovingAverage {
261 decay: initial_decay.to_f64().unwrap_or(0.9),
262 };
263
264 Self {
265 averager: WeightAverager::new(method, 1), initial_decay,
267 final_decay,
268 decay_steps: decaysteps,
269 }
270 }
271
272 pub fn update(&mut self, weights: &[Array<A, D>]) -> Result<()> {
274 let step = self.averager.step_count() as f64;
275 let progress = (step / self.decay_steps as f64).min(1.0);
276
277 let current_decay = self.initial_decay.to_f64().unwrap_or(0.9) * (1.0 - progress)
279 + self.final_decay.to_f64().unwrap_or(0.999) * progress;
280
281 self.averager
282 .set_ema_decay(try_scalar::<A, _>(current_decay)?);
283 self.averager.update(weights)
284 }
285
286 pub fn get_averaged_weights(&self) -> &[Array<A, D>] {
288 self.averager.get_averaged_weights()
289 }
290
291 pub fn initialize(&mut self, weights: &[Array<A, D>]) -> Result<()> {
293 self.averager.initialize(weights)
294 }
295}
296
297pub mod gradient_centralization {
299 use super::*;
300
301 pub fn centralize_gradients<A, D>(gradients: &mut [Array<A, D>]) -> Result<()>
303 where
304 A: Float + ScalarOperand + Debug,
305 D: Dimension,
306 {
307 for grad in gradients {
308 centralize_single_gradient(grad)?;
309 }
310 Ok(())
311 }
312
313 pub fn centralize_single_gradient<A, D>(gradient: &mut Array<A, D>) -> Result<()>
315 where
316 A: Float + ScalarOperand + Debug,
317 D: Dimension,
318 {
319 if gradient.is_empty() {
320 return Ok(());
321 }
322
323 let mean = gradient.sum() / try_scalar::<A, _>(gradient.len())?;
325
326 gradient.mapv_inplace(|x| x - mean);
328
329 Ok(())
330 }
331
332 pub fn centralize_gradients_with_scaling<A, D>(
334 gradients: &mut [Array<A, D>],
335 scale_factor: A,
336 ) -> Result<()>
337 where
338 A: Float + ScalarOperand + Debug,
339 D: Dimension,
340 {
341 centralize_gradients(gradients)?;
342
343 for grad in gradients {
345 grad.mapv_inplace(|x| x * scale_factor);
346 }
347
348 Ok(())
349 }
350}
351
352#[derive(Debug)]
354pub struct ModelEnsemble<A: Float, D: Dimension> {
355 models: Vec<Vec<Array<A, D>>>,
357 model_weights: Vec<A>,
359 ensemble_average: Option<Vec<Array<A, D>>>,
361 cache_valid: bool,
363}
364
365impl<A: Float + ScalarOperand + Debug, D: Dimension + Send + Sync> ModelEnsemble<A, D> {
366 pub fn new() -> Self {
368 Self {
369 models: Vec::new(),
370 model_weights: Vec::new(),
371 ensemble_average: None,
372 cache_valid: false,
373 }
374 }
375
376 pub fn add_model(&mut self, weights: Vec<Array<A, D>>, weight: A) -> Result<()> {
378 if !self.models.is_empty() {
379 let expected_len = self.models[0].len();
380 if weights.len() != expected_len {
381 return Err(OptimError::DimensionMismatch(format!(
382 "Expected {} weight arrays, got {}",
383 expected_len,
384 weights.len()
385 )));
386 }
387 }
388
389 self.models.push(weights);
390 self.model_weights.push(weight);
391 self.cache_valid = false;
392 Ok(())
393 }
394
395 pub fn get_ensemble_average(&mut self) -> Result<&[Array<A, D>]> {
397 if !self.cache_valid {
398 self.compute_ensemble_average()?;
399 }
400
401 self.ensemble_average
402 .as_deref()
403 .ok_or_else(|| OptimError::InvalidConfig("No models in ensemble".to_string()))
404 }
405
406 fn compute_ensemble_average(&mut self) -> Result<()> {
408 if self.models.is_empty() {
409 return Err(OptimError::InvalidConfig(
410 "No models in ensemble".to_string(),
411 ));
412 }
413
414 let total_weight: A = self.model_weights.iter().fold(A::zero(), |acc, &w| acc + w);
416 if total_weight <= A::zero() {
417 return Err(OptimError::InvalidConfig(
418 "Total ensemble weight must be > 0".to_string(),
419 ));
420 }
421
422 let num_params = self.models[0].len();
423 let mut ensemble_avg = Vec::new();
424
425 for i in 0..num_params {
427 ensemble_avg.push(Array::zeros(self.models[0][i].raw_dim()));
428 }
429
430 for (model, &weight) in self.models.iter().zip(self.model_weights.iter()) {
432 let normalized_weight = weight / total_weight;
433
434 for (avg_param, model_param) in ensemble_avg.iter_mut().zip(model.iter()) {
435 Zip::from(avg_param)
436 .and(model_param)
437 .for_each(|avg, ¶m| {
438 *avg = *avg + normalized_weight * param;
439 });
440 }
441 }
442
443 self.ensemble_average = Some(ensemble_avg);
444 self.cache_valid = true;
445 Ok(())
446 }
447
448 pub fn clear(&mut self) {
450 self.models.clear();
451 self.model_weights.clear();
452 self.ensemble_average = None;
453 self.cache_valid = false;
454 }
455
456 pub fn len(&self) -> usize {
458 self.models.len()
459 }
460
461 pub fn is_empty(&self) -> bool {
463 self.models.is_empty()
464 }
465}
466
467impl<A: Float + ScalarOperand + Debug, D: Dimension + Send + Sync> Default for ModelEnsemble<A, D> {
468 fn default() -> Self {
469 Self::new()
470 }
471}
472
473#[cfg(test)]
474mod tests {
475 use super::*;
476 use approx::assert_relative_eq;
477 use scirs2_core::ndarray::Array1;
478
479 #[test]
480 fn test_moving_average() {
481 let mut averager = WeightAverager::new(AveragingMethod::MovingAverage, 3);
482
483 let weights1 = vec![Array1::from_vec(vec![1.0, 2.0])];
484 let weights2 = vec![Array1::from_vec(vec![3.0, 4.0])];
485 let weights3 = vec![Array1::from_vec(vec![5.0, 6.0])];
486
487 averager.update(&weights1).expect("unwrap failed");
488 averager.update(&weights2).expect("unwrap failed");
489 averager.update(&weights3).expect("unwrap failed");
490
491 let avg = averager.get_averaged_weights();
492 assert!(avg[0][0] >= 1.0 && avg[0][0] <= 5.0);
495 assert!(avg[0][1] >= 2.0 && avg[0][1] <= 6.0);
496 }
497
498 #[test]
499 fn test_exponential_moving_average() {
500 let decay = 0.9;
501 let mut averager =
502 WeightAverager::new(AveragingMethod::ExponentialMovingAverage { decay }, 1);
503
504 let weights1 = vec![Array1::from_vec(vec![2.0])];
505 let weights2 = vec![Array1::from_vec(vec![4.0])];
506
507 averager.update(&weights1).expect("unwrap failed");
508 averager.update(&weights2).expect("unwrap failed");
509
510 let avg = averager.get_averaged_weights();
511 assert_relative_eq!(avg[0][0], 2.2, epsilon = 1e-6);
513 }
514
515 #[test]
516 fn test_swa() {
517 let mut averager = WeightAverager::new(AveragingMethod::StochasticWeightAveraging, 10);
518
519 let weights1 = vec![Array1::from_vec(vec![2.0])];
520 let weights2 = vec![Array1::from_vec(vec![4.0])];
521 let weights3 = vec![Array1::from_vec(vec![6.0])];
522
523 averager.update(&weights1).expect("unwrap failed"); averager.update(&weights2).expect("unwrap failed"); averager.update(&weights3).expect("unwrap failed"); let avg = averager.get_averaged_weights();
528 assert!(avg[0][0] >= 3.5 && avg[0][0] <= 5.0);
531 }
532
533 #[test]
534 fn test_gradient_centralization() {
535 let mut gradients = vec![Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0])];
536
537 gradient_centralization::centralize_gradients(&mut gradients).expect("unwrap failed");
538
539 let expected = [-1.5, -0.5, 0.5, 1.5];
542 for (actual, expected) in gradients[0].iter().zip(expected.iter()) {
543 assert_relative_eq!(*actual, *expected, epsilon = 1e-6);
544 }
545
546 let mean = gradients[0].sum() / 4.0;
548 assert_relative_eq!(mean, 0.0, epsilon = 1e-10);
549 }
550
551 #[test]
552 fn test_polyak_averager() {
553 let mut averager = PolyakAverager::new(0.5, 0.9, 10);
554
555 let weights1 = vec![Array1::from_vec(vec![2.0])];
556 let weights2 = vec![Array1::from_vec(vec![4.0])];
557
558 averager.update(&weights1).expect("unwrap failed");
559 averager.update(&weights2).expect("unwrap failed");
560
561 let avg = averager.get_averaged_weights();
562 assert!(avg[0][0] > 2.0 && avg[0][0] < 4.0); }
564
565 #[test]
566 fn test_model_ensemble() {
567 let mut ensemble = ModelEnsemble::new();
568
569 let model1 = vec![Array1::from_vec(vec![2.0, 4.0])];
570 let model2 = vec![Array1::from_vec(vec![4.0, 2.0])];
571
572 ensemble.add_model(model1, 1.0).expect("unwrap failed");
573 ensemble.add_model(model2, 1.0).expect("unwrap failed");
574
575 let avg = ensemble.get_ensemble_average().expect("unwrap failed");
576 assert_relative_eq!(avg[0][0], 3.0, epsilon = 1e-6); assert_relative_eq!(avg[0][1], 3.0, epsilon = 1e-6); }
579
580 #[test]
581 fn test_weighted_model_ensemble() {
582 let mut ensemble = ModelEnsemble::new();
583
584 let model1 = vec![Array1::from_vec(vec![2.0])];
585 let model2 = vec![Array1::from_vec(vec![4.0])];
586
587 ensemble.add_model(model1, 3.0).expect("unwrap failed"); ensemble.add_model(model2, 1.0).expect("unwrap failed"); let avg = ensemble.get_ensemble_average().expect("unwrap failed");
591 assert_relative_eq!(avg[0][0], 2.5, epsilon = 1e-6);
593 }
594
595 #[test]
596 fn test_ensemble_dimension_validation() {
597 let mut ensemble = ModelEnsemble::new();
598
599 let model1 = vec![Array1::from_vec(vec![1.0, 2.0])];
600 let model2 = vec![
601 Array1::from_vec(vec![3.0, 4.0]),
602 Array1::from_vec(vec![5.0]),
603 ]; ensemble.add_model(model1, 1.0).expect("unwrap failed");
606 assert!(ensemble.add_model(model2, 1.0).is_err());
607 }
608
609 #[test]
610 fn test_weight_averager_dimension_validation() {
611 let mut averager = WeightAverager::new(AveragingMethod::MovingAverage, 3);
612
613 let weights1 = vec![Array1::from_vec(vec![1.0, 2.0])];
614 let weights2 = vec![
615 Array1::from_vec(vec![3.0, 4.0]),
616 Array1::from_vec(vec![5.0]),
617 ]; averager.update(&weights1).expect("unwrap failed");
620 assert!(averager.update(&weights2).is_err());
621 }
622
623 #[test]
624 fn test_gradient_centralization_with_scaling() {
625 let mut gradients = vec![Array1::from_vec(vec![1.0, 3.0])]; gradient_centralization::centralize_gradients_with_scaling(&mut gradients, 2.0)
628 .expect("unwrap failed");
629
630 assert_relative_eq!(gradients[0][0], -2.0, epsilon = 1e-6);
632 assert_relative_eq!(gradients[0][1], 2.0, epsilon = 1e-6);
633 }
634}