1use crate::{Optimizer, OptimizerError, OptimizerResult, OptimizerState};
46use parking_lot::RwLock;
47use std::collections::HashMap;
48use std::sync::Arc;
49use torsh_tensor::Tensor;
50
51#[derive(Debug, Clone)]
57pub struct EWCConfig {
58 pub importance: f32,
60 pub fisher_sample_size: usize,
62 pub diagonal_fisher: bool,
64}
65
66impl Default for EWCConfig {
67 fn default() -> Self {
68 Self {
69 importance: 1000.0,
70 fisher_sample_size: 200,
71 diagonal_fisher: true,
72 }
73 }
74}
75
76pub struct EWCOptimizer<O: Optimizer> {
81 base_optimizer: O,
83 config: EWCConfig,
85 fisher_information: HashMap<String, Tensor>,
87 optimal_params: HashMap<String, Tensor>,
89 current_task: usize,
91 param_groups: Vec<Arc<RwLock<Tensor>>>,
93}
94
95impl<O: Optimizer> EWCOptimizer<O> {
96 pub fn new(
98 base_optimizer: O,
99 params: Vec<Arc<RwLock<Tensor>>>,
100 config: EWCConfig,
101 ) -> OptimizerResult<Self> {
102 Ok(Self {
103 base_optimizer,
104 config,
105 fisher_information: HashMap::new(),
106 optimal_params: HashMap::new(),
107 current_task: 0,
108 param_groups: params,
109 })
110 }
111
112 pub fn with_defaults(
114 base_optimizer: O,
115 params: Vec<Arc<RwLock<Tensor>>>,
116 ) -> OptimizerResult<Self> {
117 Self::new(base_optimizer, params, EWCConfig::default())
118 }
119
120 pub fn consolidate_task(&mut self) -> OptimizerResult<()> {
122 for (i, param) in self.param_groups.iter().enumerate() {
124 let param_key = format!("param_{}", i);
125 let param_read = param.read();
126 self.optimal_params
127 .insert(param_key.clone(), param_read.clone());
128 }
129
130 self.compute_fisher_diagonal()?;
132
133 self.current_task += 1;
134 Ok(())
135 }
136
137 fn compute_fisher_diagonal(&mut self) -> OptimizerResult<()> {
139 for (i, param) in self.param_groups.iter().enumerate() {
141 let param_key = format!("param_{}", i);
142
143 let param_read = param.read();
145 if let Some(grad) = param_read.grad() {
146 let fisher = grad.mul(&grad)?; if let Some(existing_fisher) = self.fisher_information.get(¶m_key) {
150 let accumulated = existing_fisher.add(&fisher)?;
151 self.fisher_information.insert(param_key, accumulated);
152 } else {
153 self.fisher_information.insert(param_key, fisher);
154 }
155 }
156 }
157
158 Ok(())
159 }
160
161 fn apply_ewc_penalty(&mut self) -> OptimizerResult<()> {
163 if self.optimal_params.is_empty() {
164 return Ok(());
166 }
167
168 for (i, param) in self.param_groups.iter().enumerate() {
169 let param_key = format!("param_{}", i);
170
171 if let (Some(fisher), Some(optimal)) = (
172 self.fisher_information.get(¶m_key),
173 self.optimal_params.get(¶m_key),
174 ) {
175 let mut param_write = param.write();
176
177 let diff = param_write.sub(optimal)?;
179 let penalty = fisher.mul(&diff)?;
180 let scaled_penalty = penalty.mul_scalar(self.config.importance)?;
181
182 if let Some(grad) = param_write.grad() {
184 let new_grad = grad.add(&scaled_penalty)?;
185 param_write.set_grad(Some(new_grad));
186 }
187 }
188 }
189
190 Ok(())
191 }
192}
193
194impl<O: Optimizer> Optimizer for EWCOptimizer<O> {
195 fn step(&mut self) -> OptimizerResult<()> {
196 self.apply_ewc_penalty()?;
198
199 self.base_optimizer.step()
201 }
202
203 fn zero_grad(&mut self) {
204 self.base_optimizer.zero_grad();
205 }
206
207 fn get_lr(&self) -> Vec<f32> {
208 self.base_optimizer.get_lr()
209 }
210
211 fn set_lr(&mut self, lr: f32) {
212 self.base_optimizer.set_lr(lr);
213 }
214
215 fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
216 self.base_optimizer.add_param_group(params, options);
217 }
218
219 fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
220 self.param_groups.clone()
221 }
222
223 fn state_dict(&self) -> OptimizerResult<OptimizerState> {
224 let mut state = self.base_optimizer.state_dict()?;
225 state.optimizer_type = format!("EWC({})", state.optimizer_type);
226 state
227 .global_state
228 .insert("current_task".to_string(), self.current_task as f32);
229 state
230 .global_state
231 .insert("importance".to_string(), self.config.importance);
232 Ok(state)
233 }
234
235 fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
236 self.base_optimizer.load_state_dict(state)
237 }
238}
239
240#[derive(Debug, Clone)]
246pub struct SIConfig {
247 pub damping: f32,
249 pub importance: f32,
251}
252
253impl Default for SIConfig {
254 fn default() -> Self {
255 Self {
256 damping: 0.1,
257 importance: 1.0,
258 }
259 }
260}
261
262pub struct SIOptimizer<O: Optimizer> {
267 base_optimizer: O,
269 config: SIConfig,
271 path_integral: HashMap<String, Tensor>,
273 prev_params: HashMap<String, Tensor>,
275 importance: HashMap<String, Tensor>,
277 current_task: usize,
279 param_groups: Vec<Arc<RwLock<Tensor>>>,
281}
282
283impl<O: Optimizer> SIOptimizer<O> {
284 pub fn new(
286 base_optimizer: O,
287 params: Vec<Arc<RwLock<Tensor>>>,
288 config: SIConfig,
289 ) -> OptimizerResult<Self> {
290 let mut prev_params = HashMap::new();
291 let mut path_integral = HashMap::new();
292
293 for (i, param) in params.iter().enumerate() {
295 let param_key = format!("param_{}", i);
296 let param_read = param.read();
297 prev_params.insert(param_key.clone(), param_read.clone());
298
299 let shape_owned = param_read.shape().dims().to_vec();
300 drop(param_read);
301 let zeros = torsh_tensor::creation::zeros(&shape_owned)?;
302 path_integral.insert(param_key, zeros);
303 }
304
305 Ok(Self {
306 base_optimizer,
307 config,
308 path_integral,
309 prev_params,
310 importance: HashMap::new(),
311 current_task: 0,
312 param_groups: params,
313 })
314 }
315
316 pub fn with_defaults(
318 base_optimizer: O,
319 params: Vec<Arc<RwLock<Tensor>>>,
320 ) -> OptimizerResult<Self> {
321 Self::new(base_optimizer, params, SIConfig::default())
322 }
323
324 fn update_path_integral(&mut self) -> OptimizerResult<()> {
326 for (i, param) in self.param_groups.iter().enumerate() {
327 let param_key = format!("param_{}", i);
328
329 let param_read = param.read();
330 if let Some(grad) = param_read.grad() {
331 if let Some(prev_param) = self.prev_params.get(¶m_key) {
332 let delta = param_read.sub(prev_param)?;
334
335 let contribution = grad.mul(&delta)?;
337 let neg_contribution = contribution.mul_scalar(-1.0)?;
338
339 if let Some(omega) = self.path_integral.get_mut(¶m_key) {
340 *omega = omega.add(&neg_contribution)?;
341 }
342
343 self.prev_params.insert(param_key, param_read.clone());
345 }
346 }
347 }
348
349 Ok(())
350 }
351
352 pub fn consolidate_task(&mut self) -> OptimizerResult<()> {
354 for (i, param) in self.param_groups.iter().enumerate() {
355 let param_key = format!("param_{}", i);
356
357 if let Some(omega) = self.path_integral.get(¶m_key) {
358 if let Some(prev_param) = self.prev_params.get(¶m_key) {
360 let param_read = param.read();
361 let delta = param_read.sub(prev_param)?;
362 let delta_sq = delta.mul(&delta)?;
363 let denom = delta_sq.add_scalar(self.config.damping)?;
364
365 let task_importance = omega.div(&denom)?;
366
367 if let Some(existing) = self.importance.get(¶m_key) {
369 let accumulated = existing.add(&task_importance)?;
370 self.importance.insert(param_key.clone(), accumulated);
371 } else {
372 self.importance.insert(param_key.clone(), task_importance);
373 }
374
375 let shape_owned = param_read.shape().dims().to_vec();
377 let zeros = torsh_tensor::creation::zeros(&shape_owned)?;
378 self.path_integral.insert(param_key, zeros);
379 }
380 }
381 }
382
383 self.current_task += 1;
384 Ok(())
385 }
386
387 fn apply_si_penalty(&mut self) -> OptimizerResult<()> {
389 if self.importance.is_empty() {
390 return Ok(());
391 }
392
393 for (i, param) in self.param_groups.iter().enumerate() {
394 let param_key = format!("param_{}", i);
395
396 if let (Some(importance), Some(prev_param)) = (
397 self.importance.get(¶m_key),
398 self.prev_params.get(¶m_key),
399 ) {
400 let mut param_write = param.write();
401
402 let diff = param_write.sub(prev_param)?;
404 let penalty = importance.mul(&diff)?;
405 let scaled_penalty = penalty.mul_scalar(self.config.importance)?;
406
407 if let Some(grad) = param_write.grad() {
408 let new_grad = grad.add(&scaled_penalty)?;
409 param_write.set_grad(Some(new_grad));
410 }
411 }
412 }
413
414 Ok(())
415 }
416}
417
418impl<O: Optimizer> Optimizer for SIOptimizer<O> {
419 fn step(&mut self) -> OptimizerResult<()> {
420 self.update_path_integral()?;
422
423 self.apply_si_penalty()?;
425
426 self.base_optimizer.step()
428 }
429
430 fn zero_grad(&mut self) {
431 self.base_optimizer.zero_grad();
432 }
433
434 fn get_lr(&self) -> Vec<f32> {
435 self.base_optimizer.get_lr()
436 }
437
438 fn set_lr(&mut self, lr: f32) {
439 self.base_optimizer.set_lr(lr);
440 }
441
442 fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
443 self.base_optimizer.add_param_group(params, options);
444 }
445
446 fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
447 self.param_groups.clone()
448 }
449
450 fn state_dict(&self) -> OptimizerResult<OptimizerState> {
451 let mut state = self.base_optimizer.state_dict()?;
452 state.optimizer_type = format!("SI({})", state.optimizer_type);
453 state
454 .global_state
455 .insert("current_task".to_string(), self.current_task as f32);
456 Ok(state)
457 }
458
459 fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
460 self.base_optimizer.load_state_dict(state)
461 }
462}
463
464#[derive(Debug, Clone)]
470pub struct MASConfig {
471 pub importance: f32,
473 pub n_samples: usize,
475}
476
477impl Default for MASConfig {
478 fn default() -> Self {
479 Self {
480 importance: 1.0,
481 n_samples: 100,
482 }
483 }
484}
485
486pub struct MASOptimizer<O: Optimizer> {
491 base_optimizer: O,
493 config: MASConfig,
495 importance: HashMap<String, Tensor>,
497 optimal_params: HashMap<String, Tensor>,
499 current_task: usize,
501 param_groups: Vec<Arc<RwLock<Tensor>>>,
503}
504
505impl<O: Optimizer> MASOptimizer<O> {
506 pub fn new(
508 base_optimizer: O,
509 params: Vec<Arc<RwLock<Tensor>>>,
510 config: MASConfig,
511 ) -> OptimizerResult<Self> {
512 Ok(Self {
513 base_optimizer,
514 config,
515 importance: HashMap::new(),
516 optimal_params: HashMap::new(),
517 current_task: 0,
518 param_groups: params,
519 })
520 }
521
522 pub fn with_defaults(
524 base_optimizer: O,
525 params: Vec<Arc<RwLock<Tensor>>>,
526 ) -> OptimizerResult<Self> {
527 Self::new(base_optimizer, params, MASConfig::default())
528 }
529
530 pub fn compute_importance(&mut self) -> OptimizerResult<()> {
532 for (i, param) in self.param_groups.iter().enumerate() {
534 let param_key = format!("param_{}", i);
535
536 let param_read = param.read();
537 if let Some(grad) = param_read.grad() {
538 let grad_abs = grad.abs()?;
540
541 if let Some(existing) = self.importance.get(¶m_key) {
542 let accumulated = existing.add(&grad_abs)?;
543 self.importance.insert(param_key, accumulated);
544 } else {
545 self.importance.insert(param_key, grad_abs);
546 }
547 }
548 }
549
550 Ok(())
551 }
552
553 pub fn consolidate_task(&mut self) -> OptimizerResult<()> {
555 for (i, param) in self.param_groups.iter().enumerate() {
557 let param_key = format!("param_{}", i);
558 let param_read = param.read();
559 self.optimal_params.insert(param_key, param_read.clone());
560 }
561
562 self.current_task += 1;
563 Ok(())
564 }
565
566 fn apply_mas_penalty(&mut self) -> OptimizerResult<()> {
568 if self.importance.is_empty() {
569 return Ok(());
570 }
571
572 for (i, param) in self.param_groups.iter().enumerate() {
573 let param_key = format!("param_{}", i);
574
575 if let (Some(importance), Some(optimal)) = (
576 self.importance.get(¶m_key),
577 self.optimal_params.get(¶m_key),
578 ) {
579 let mut param_write = param.write();
580
581 let diff = param_write.sub(optimal)?;
583 let penalty = importance.mul(&diff)?;
584 let scaled_penalty = penalty.mul_scalar(self.config.importance)?;
585
586 if let Some(grad) = param_write.grad() {
587 let new_grad = grad.add(&scaled_penalty)?;
588 param_write.set_grad(Some(new_grad));
589 }
590 }
591 }
592
593 Ok(())
594 }
595}
596
597impl<O: Optimizer> Optimizer for MASOptimizer<O> {
598 fn step(&mut self) -> OptimizerResult<()> {
599 self.apply_mas_penalty()?;
601
602 self.base_optimizer.step()
604 }
605
606 fn zero_grad(&mut self) {
607 self.base_optimizer.zero_grad();
608 }
609
610 fn get_lr(&self) -> Vec<f32> {
611 self.base_optimizer.get_lr()
612 }
613
614 fn set_lr(&mut self, lr: f32) {
615 self.base_optimizer.set_lr(lr);
616 }
617
618 fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
619 self.base_optimizer.add_param_group(params, options);
620 }
621
622 fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
623 self.param_groups.clone()
624 }
625
626 fn state_dict(&self) -> OptimizerResult<OptimizerState> {
627 let mut state = self.base_optimizer.state_dict()?;
628 state.optimizer_type = format!("MAS({})", state.optimizer_type);
629 state
630 .global_state
631 .insert("current_task".to_string(), self.current_task as f32);
632 Ok(state)
633 }
634
635 fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
636 self.base_optimizer.load_state_dict(state)
637 }
638}
639
640#[cfg(test)]
645mod tests {
646 use super::*;
647 use crate::sgd::SGD;
648 use torsh_tensor::creation::randn;
649
650 #[test]
651 fn test_ewc_config_default() {
652 let config = EWCConfig::default();
653 assert_eq!(config.importance, 1000.0);
654 assert_eq!(config.fisher_sample_size, 200);
655 assert!(config.diagonal_fisher);
656 }
657
658 #[test]
659 fn test_ewc_optimizer_creation() -> OptimizerResult<()> {
660 let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
661 let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
662
663 let optimizer = EWCOptimizer::with_defaults(base, vec![param])?;
664 assert_eq!(optimizer.current_task, 0);
665 Ok(())
666 }
667
668 #[test]
669 fn test_ewc_consolidate_task() -> OptimizerResult<()> {
670 let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
671
672 {
674 let mut p = param.write();
675 let grad = randn::<f32>(&[5, 5])?;
676 p.set_grad(Some(grad));
677 }
678
679 let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
680 let mut optimizer = EWCOptimizer::with_defaults(base, vec![param])?;
681
682 optimizer.consolidate_task()?;
683 assert_eq!(optimizer.current_task, 1);
684 assert!(!optimizer.optimal_params.is_empty());
685 Ok(())
686 }
687
688 #[test]
689 fn test_si_config_default() {
690 let config = SIConfig::default();
691 assert_eq!(config.damping, 0.1);
692 assert_eq!(config.importance, 1.0);
693 }
694
695 #[test]
696 fn test_si_optimizer_creation() -> OptimizerResult<()> {
697 let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
698 let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
699
700 let optimizer = SIOptimizer::with_defaults(base, vec![param])?;
701 assert_eq!(optimizer.current_task, 0);
702 Ok(())
703 }
704
705 #[test]
706 fn test_si_consolidate_task() -> OptimizerResult<()> {
707 let param = Arc::new(RwLock::new(randn::<f32>(&[3, 3])?));
708
709 {
711 let mut p = param.write();
712 let grad = randn::<f32>(&[3, 3])?;
713 p.set_grad(Some(grad));
714 }
715
716 let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
717 let mut optimizer = SIOptimizer::with_defaults(base, vec![param])?;
718
719 optimizer.consolidate_task()?;
720 assert_eq!(optimizer.current_task, 1);
721 Ok(())
722 }
723
724 #[test]
725 fn test_mas_config_default() {
726 let config = MASConfig::default();
727 assert_eq!(config.importance, 1.0);
728 assert_eq!(config.n_samples, 100);
729 }
730
731 #[test]
732 fn test_mas_optimizer_creation() -> OptimizerResult<()> {
733 let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
734 let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
735
736 let optimizer = MASOptimizer::with_defaults(base, vec![param])?;
737 assert_eq!(optimizer.current_task, 0);
738 Ok(())
739 }
740
741 #[test]
742 fn test_mas_compute_importance() -> OptimizerResult<()> {
743 let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
744
745 {
747 let mut p = param.write();
748 let grad = randn::<f32>(&[5, 5])?;
749 p.set_grad(Some(grad));
750 }
751
752 let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
753 let mut optimizer = MASOptimizer::with_defaults(base, vec![param])?;
754
755 optimizer.compute_importance()?;
756 assert!(!optimizer.importance.is_empty());
757 Ok(())
758 }
759
760 #[test]
761 fn test_ewc_step() -> OptimizerResult<()> {
762 let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
763
764 {
766 let mut p = param.write();
767 let grad = randn::<f32>(&[2, 2])?;
768 p.set_grad(Some(grad));
769 }
770
771 let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
772 let mut optimizer = EWCOptimizer::with_defaults(base, vec![param])?;
773
774 optimizer.step()?;
776 Ok(())
777 }
778}