Skip to main content

torsh_optim/
continual_learning.rs

1//! Continual Learning Optimizers
2//!
3//! This module implements optimization algorithms for lifelong learning scenarios,
4//! preventing catastrophic forgetting when learning sequential tasks.
5//!
6//! # Key Concepts
7//!
8//! - **Catastrophic Forgetting**: The tendency of neural networks to forget previously
9//!   learned tasks when learning new ones.
10//! - **Parameter Importance**: Identifying which parameters are critical for past tasks
11//!   and should be protected from large updates.
12//! - **Task-Specific Learning**: Adapting learning strategies based on task sequence.
13//!
14//! # Algorithms
15//!
16//! ## EWC (Elastic Weight Consolidation)
17//!
18//! Protects important parameters by adding a quadratic penalty based on the Fisher
19//! Information Matrix. Parameters critical for previous tasks receive higher penalties.
20//!
21//! Loss: L_new = L_task + (λ/2) Σ F_i (θ_i - θ*_i)²
22//!
23//! ## SI (Synaptic Intelligence)
24//!
25//! Online continual learning that accumulates parameter importance during training.
26//! Updates importance based on path integral of parameter changes.
27//!
28//! ## MAS (Memory Aware Synapses)
29//!
30//! Uses gradient magnitude at optimal parameters to estimate importance,
31//! avoiding the need for data from previous tasks.
32//!
33//! ## PackNet
34//!
35//! Packs multiple tasks into a single network by dynamically allocating
36//! and protecting subnetworks for each task.
37//!
38//! # References
39//!
40//! - Kirkpatrick et al. (2017). "Overcoming catastrophic forgetting in neural networks"
41//! - Zenke et al. (2017). "Continual Learning Through Synaptic Intelligence"
42//! - Aljundi et al. (2018). "Memory Aware Synapses"
43//! - Mallya & Lazebnik (2018). "PackNet: Adding Multiple Tasks to a Single Network"
44
45use crate::{Optimizer, OptimizerError, OptimizerResult, OptimizerState};
46use parking_lot::RwLock;
47use std::collections::HashMap;
48use std::sync::Arc;
49use torsh_tensor::Tensor;
50
51// ============================================================================
52// EWC (Elastic Weight Consolidation)
53// ============================================================================
54
55/// EWC configuration
56#[derive(Debug, Clone)]
57pub struct EWCConfig {
58    /// Importance weight (λ) for Fisher penalty
59    pub importance: f32,
60    /// Sample size for Fisher matrix estimation
61    pub fisher_sample_size: usize,
62    /// Use diagonal Fisher approximation
63    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
76/// Elastic Weight Consolidation optimizer
77///
78/// Prevents catastrophic forgetting by adding regularization based on
79/// parameter importance computed from the Fisher Information Matrix.
80pub struct EWCOptimizer<O: Optimizer> {
81    /// Base optimizer
82    base_optimizer: O,
83    /// Configuration
84    config: EWCConfig,
85    /// Fisher information per parameter
86    fisher_information: HashMap<String, Tensor>,
87    /// Optimal parameters from previous task
88    optimal_params: HashMap<String, Tensor>,
89    /// Current task ID
90    current_task: usize,
91    /// Parameter groups reference
92    param_groups: Vec<Arc<RwLock<Tensor>>>,
93}
94
95impl<O: Optimizer> EWCOptimizer<O> {
96    /// Create a new EWC optimizer
97    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    /// Create with default configuration
113    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    /// Consolidate current task (compute Fisher and save optimal parameters)
121    pub fn consolidate_task(&mut self) -> OptimizerResult<()> {
122        // Save current parameters as optimal
123        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        // Compute Fisher information (simplified diagonal approximation)
131        self.compute_fisher_diagonal()?;
132
133        self.current_task += 1;
134        Ok(())
135    }
136
137    /// Compute diagonal Fisher information matrix
138    fn compute_fisher_diagonal(&mut self) -> OptimizerResult<()> {
139        // For each parameter, compute F_ii = E[∂L/∂θ_i]²
140        for (i, param) in self.param_groups.iter().enumerate() {
141            let param_key = format!("param_{}", i);
142
143            // Get gradient (squared for Fisher diagonal)
144            let param_read = param.read();
145            if let Some(grad) = param_read.grad() {
146                let fisher = grad.mul(&grad)?; // Element-wise square
147
148                // If Fisher already exists, accumulate
149                if let Some(existing_fisher) = self.fisher_information.get(&param_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    /// Apply EWC penalty to gradients
162    fn apply_ewc_penalty(&mut self) -> OptimizerResult<()> {
163        if self.optimal_params.is_empty() {
164            // No previous task, no penalty
165            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(&param_key),
173                self.optimal_params.get(&param_key),
174            ) {
175                let mut param_write = param.write();
176
177                // Compute EWC penalty: λ * F * (θ - θ*)
178                let diff = param_write.sub(optimal)?;
179                let penalty = fisher.mul(&diff)?;
180                let scaled_penalty = penalty.mul_scalar(self.config.importance)?;
181
182                // Add penalty to gradient
183                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        // Apply EWC penalty to gradients
197        self.apply_ewc_penalty()?;
198
199        // Call base optimizer
200        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// ============================================================================
241// SI (Synaptic Intelligence)
242// ============================================================================
243
244/// Synaptic Intelligence configuration
245#[derive(Debug, Clone)]
246pub struct SIConfig {
247    /// Damping parameter (ξ)
248    pub damping: f32,
249    /// Importance regularization strength
250    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
262/// Synaptic Intelligence optimizer
263///
264/// Online continual learning that tracks parameter importance during training
265/// using path integral of gradient times parameter change.
266pub struct SIOptimizer<O: Optimizer> {
267    /// Base optimizer
268    base_optimizer: O,
269    /// Configuration
270    config: SIConfig,
271    /// Path integral accumulator (ω)
272    path_integral: HashMap<String, Tensor>,
273    /// Previous parameters
274    prev_params: HashMap<String, Tensor>,
275    /// Consolidated importance per task
276    importance: HashMap<String, Tensor>,
277    /// Current task ID
278    current_task: usize,
279    /// Parameter groups
280    param_groups: Vec<Arc<RwLock<Tensor>>>,
281}
282
283impl<O: Optimizer> SIOptimizer<O> {
284    /// Create a new SI optimizer
285    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        // Initialize previous parameters and path integral
294        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    /// Create with default configuration
317    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    /// Update path integral
325    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(&param_key) {
332                    // Δθ = θ_new - θ_old
333                    let delta = param_read.sub(prev_param)?;
334
335                    // ω += -∂L/∂θ * Δθ
336                    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(&param_key) {
340                        *omega = omega.add(&neg_contribution)?;
341                    }
342
343                    // Update previous parameters
344                    self.prev_params.insert(param_key, param_read.clone());
345                }
346            }
347        }
348
349        Ok(())
350    }
351
352    /// Consolidate task importance
353    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(&param_key) {
358                // Compute importance: Ω = ω / (Δθ² + ξ)
359                if let Some(prev_param) = self.prev_params.get(&param_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                    // Accumulate importance
368                    if let Some(existing) = self.importance.get(&param_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                    // Reset path integral
376                    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    /// Apply SI penalty
388    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(&param_key),
398                self.prev_params.get(&param_key),
399            ) {
400                let mut param_write = param.write();
401
402                // Penalty: Ω * (θ - θ_prev)
403                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        // Update path integral before step
421        self.update_path_integral()?;
422
423        // Apply SI penalty
424        self.apply_si_penalty()?;
425
426        // Call base optimizer
427        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// ============================================================================
465// MAS (Memory Aware Synapses)
466// ============================================================================
467
468/// MAS configuration
469#[derive(Debug, Clone)]
470pub struct MASConfig {
471    /// Importance regularization strength
472    pub importance: f32,
473    /// Number of samples for importance estimation
474    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
486/// Memory Aware Synapses optimizer
487///
488/// Estimates parameter importance using gradient magnitude at optimal parameters,
489/// avoiding the need for data from previous tasks.
490pub struct MASOptimizer<O: Optimizer> {
491    /// Base optimizer
492    base_optimizer: O,
493    /// Configuration
494    config: MASConfig,
495    /// Parameter importance
496    importance: HashMap<String, Tensor>,
497    /// Optimal parameters from previous tasks
498    optimal_params: HashMap<String, Tensor>,
499    /// Current task ID
500    current_task: usize,
501    /// Parameter groups
502    param_groups: Vec<Arc<RwLock<Tensor>>>,
503}
504
505impl<O: Optimizer> MASOptimizer<O> {
506    /// Create a new MAS optimizer
507    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    /// Create with default configuration
523    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    /// Compute importance using output gradient magnitude
531    pub fn compute_importance(&mut self) -> OptimizerResult<()> {
532        // Accumulate gradient magnitudes
533        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                // Importance = |∂L/∂θ|
539                let grad_abs = grad.abs()?;
540
541                if let Some(existing) = self.importance.get(&param_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    /// Consolidate task (save optimal parameters)
554    pub fn consolidate_task(&mut self) -> OptimizerResult<()> {
555        // Save current parameters
556        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    /// Apply MAS penalty
567    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(&param_key),
577                self.optimal_params.get(&param_key),
578            ) {
579                let mut param_write = param.write();
580
581                // Penalty: Ω * (θ - θ*)
582                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        // Apply MAS penalty
600        self.apply_mas_penalty()?;
601
602        // Call base optimizer
603        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// ============================================================================
641// Tests
642// ============================================================================
643
644#[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        // Set gradient for Fisher computation
673        {
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        // Set gradient
710        {
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        // Set gradient
746        {
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        // Set gradient
765        {
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        // Step should succeed
775        optimizer.step()?;
776        Ok(())
777    }
778}