Skip to main content

optirs_core/optimizers/
grouped_adam.rs

1// Adam optimizer with parameter group support
2
3use crate::error::{OptimError, Result};
4use crate::optimizers::Optimizer;
5use crate::parameter_groups::{
6    GroupManager, GroupedOptimizer, ParameterGroup, ParameterGroupConfig,
7};
8use scirs2_core::ndarray::{Array, Dimension, ScalarOperand};
9use scirs2_core::numeric::Float;
10use std::collections::HashMap;
11use std::fmt::Debug;
12
13/// Adam optimizer with parameter group support
14///
15/// This optimizer allows different parameter groups to have different
16/// hyperparameters (learning rate, weight decay, betas).
17///
18/// # Example
19///
20/// ```no_run
21/// use scirs2_core::ndarray::Array1;
22/// use optirs_core::optimizers::{GroupedAdam, Optimizer};
23/// use optirs_core::parameter_groups::{GroupedOptimizer, ParameterGroupConfig};
24///
25/// // Create grouped optimizer
26/// let mut optimizer = GroupedAdam::new(0.001);
27///
28/// // Add parameter groups with different learning rates
29/// let params_fast = vec![Array1::zeros(5)];
30/// let config_fast = ParameterGroupConfig::new().with_learning_rate(0.01);
31/// let group_fast = optimizer.add_group(params_fast, config_fast).expect("optimizer.add_group succeeds");
32///
33/// let params_slow = vec![Array1::zeros(3)];
34/// let config_slow = ParameterGroupConfig::new().with_learning_rate(0.0001);
35/// let group_slow = optimizer.add_group(params_slow, config_slow).expect("optimizer.add_group succeeds");
36///
37/// // Optimize each group separately
38/// let grads_fast = vec![Array1::ones(5)];
39/// let updated_fast = optimizer.step_group(group_fast, &grads_fast).expect("optimizer.step_group succeeds");
40///
41/// let grads_slow = vec![Array1::ones(3)];
42/// let updated_slow = optimizer.step_group(group_slow, &grads_slow).expect("optimizer.step_group succeeds");
43/// ```
44#[derive(Debug)]
45pub struct GroupedAdam<A: Float + Send + Sync, D: Dimension> {
46    /// Default learning rate
47    defaultlr: A,
48    /// Default beta1
49    default_beta1: A,
50    /// Default beta2
51    default_beta2: A,
52    /// Default weight decay
53    default_weight_decay: A,
54    /// Epsilon to prevent division by zero
55    epsilon: A,
56    /// AMSGrad flag
57    amsgrad: bool,
58    /// Parameter groups
59    group_manager: GroupManager<A, D>,
60    /// Global step counter (total number of group updates performed)
61    step: usize,
62    /// Per-group step counters, used for per-group bias correction
63    group_steps: HashMap<usize, usize>,
64    /// Group used by the plain [`Optimizer::step`] entry point
65    ///
66    /// Cached so that repeated `step` calls reuse one group instead of appending a
67    /// fresh (state-free) group on every call.
68    implicit_group: Option<usize>,
69}
70
71impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync> GroupedAdam<A, D> {
72    /// Create a new grouped Adam optimizer
73    pub fn new(defaultlr: A) -> Self {
74        Self {
75            defaultlr,
76            default_beta1: A::from(0.9).expect("GroupedAdam: default beta1 (0.9) must fit in A"),
77            default_beta2: A::from(0.999)
78                .expect("GroupedAdam: default beta2 (0.999) must fit in A"),
79            default_weight_decay: A::zero(),
80            epsilon: A::from(1e-8).expect("GroupedAdam: default epsilon (1e-8) must fit in A"),
81            amsgrad: false,
82            group_manager: GroupManager::new(),
83            step: 0,
84            group_steps: HashMap::new(),
85            implicit_group: None,
86        }
87    }
88
89    /// Returns the number of update steps applied to `groupid`
90    pub fn group_step_count(&self, groupid: usize) -> usize {
91        self.group_steps.get(&groupid).copied().unwrap_or(0)
92    }
93
94    /// Removes every parameter group and all associated state
95    pub fn clear_groups(&mut self) {
96        self.group_manager = GroupManager::new();
97        self.group_steps.clear();
98        self.implicit_group = None;
99        self.step = 0;
100    }
101
102    /// Set default beta1
103    pub fn with_beta1(mut self, beta1: A) -> Self {
104        self.default_beta1 = beta1;
105        self
106    }
107
108    /// Set default beta2
109    pub fn with_beta2(mut self, beta2: A) -> Self {
110        self.default_beta2 = beta2;
111        self
112    }
113
114    /// Set default weight decay
115    pub fn with_weight_decay(mut self, weight_decay: A) -> Self {
116        self.default_weight_decay = weight_decay;
117        self
118    }
119
120    /// Enable AMSGrad
121    pub fn with_amsgrad(mut self) -> Self {
122        self.amsgrad = true;
123        self
124    }
125
126    /// Initialize state for a group
127    fn init_group_state(&mut self, groupid: usize) -> Result<()> {
128        let group = self.group_manager.get_group_mut(groupid)?;
129
130        if group.state.is_empty() {
131            let mut m_t = Vec::new();
132            let mut v_t = Vec::new();
133            let mut v_hat_max = Vec::new();
134
135            for param in &group.params {
136                m_t.push(Array::zeros(param.raw_dim()));
137                v_t.push(Array::zeros(param.raw_dim()));
138                if self.amsgrad {
139                    v_hat_max.push(Array::zeros(param.raw_dim()));
140                }
141            }
142
143            group.state.insert("m_t".to_string(), m_t);
144            group.state.insert("v_t".to_string(), v_t);
145            if self.amsgrad {
146                group.state.insert("v_hat_max".to_string(), v_hat_max);
147            }
148        }
149
150        Ok(())
151    }
152
153    /// Step for a specific group
154    ///
155    /// `group_step` is the 1-based update count of this group and drives bias
156    /// correction, so each group is bias-corrected against its own history.
157    fn step_group_internal(
158        &mut self,
159        groupid: usize,
160        group_step: usize,
161        gradients: &[Array<A, D>],
162    ) -> Result<Vec<Array<A, D>>> {
163        let t = i32::try_from(group_step).map_err(|_| {
164            OptimError::InvalidConfig(
165                "Timestep too large for bias correction calculation".to_string(),
166            )
167        })?;
168
169        // Initialize state if needed
170        self.init_group_state(groupid)?;
171
172        let group = self.group_manager.get_group_mut(groupid)?;
173
174        if gradients.len() != group.params.len() {
175            return Err(OptimError::InvalidConfig(format!(
176                "Number of gradients ({}) doesn't match number of parameters ({})",
177                gradients.len(),
178                group.params.len()
179            )));
180        }
181
182        // Get hyperparameters for this group
183        let lr = group.learning_rate(self.defaultlr);
184        let beta1 = group.get_custom_param("beta1", self.default_beta1);
185        let beta2 = group.get_custom_param("beta2", self.default_beta2);
186        let weightdecay = group.weight_decay(self.default_weight_decay);
187
188        let mut updated_params = Vec::new();
189
190        // Process each parameter
191        for i in 0..group.params.len() {
192            let param = &group.params[i];
193            let grad = &gradients[i];
194
195            // Apply weight decay
196            let grad_with_decay = if weightdecay > A::zero() {
197                grad + &(param * weightdecay)
198            } else {
199                grad.clone()
200            };
201
202            // Update states and compute new parameters
203            let updated = {
204                // Update first moment
205                let m_t = group.state.get_mut("m_t").ok_or_else(|| {
206                    OptimError::InvalidConfig("missing 'm_t' state for group".to_string())
207                })?;
208                m_t[i] = &m_t[i] * beta1 + &grad_with_decay * (A::one() - beta1);
209                let m_hat = &m_t[i] / (A::one() - beta1.powi(t));
210
211                // Update second moment
212                let v_t = group.state.get_mut("v_t").ok_or_else(|| {
213                    OptimError::InvalidConfig("missing 'v_t' state for group".to_string())
214                })?;
215                v_t[i] = &v_t[i] * beta2 + &grad_with_decay * &grad_with_decay * (A::one() - beta2);
216                let v_hat = &v_t[i] / (A::one() - beta2.powi(t));
217
218                // Update parameters
219                if self.amsgrad {
220                    let v_hat_max = group.state.get_mut("v_hat_max").ok_or_else(|| {
221                        OptimError::InvalidConfig("missing 'v_hat_max' state for group".to_string())
222                    })?;
223                    v_hat_max[i].zip_mut_with(&v_hat, |a, &b| *a = a.max(b));
224                    param - &(&m_hat * lr / (&v_hat_max[i].mapv(|x| x.sqrt()) + self.epsilon))
225                } else {
226                    param - &(&m_hat * lr / (&v_hat.mapv(|x| x.sqrt()) + self.epsilon))
227                }
228            };
229
230            updated_params.push(updated);
231        }
232
233        // Update group parameters
234        group.params = updated_params.clone();
235
236        Ok(updated_params)
237    }
238}
239
240impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync>
241    GroupedOptimizer<A, D> for GroupedAdam<A, D>
242{
243    fn add_group(
244        &mut self,
245        params: Vec<Array<A, D>>,
246        config: ParameterGroupConfig<A>,
247    ) -> Result<usize> {
248        Ok(self.group_manager.add_group(params, config))
249    }
250
251    fn get_group(&self, groupid: usize) -> Result<&ParameterGroup<A, D>> {
252        self.group_manager.get_group(groupid)
253    }
254
255    fn get_group_mut(&mut self, groupid: usize) -> Result<&mut ParameterGroup<A, D>> {
256        self.group_manager.get_group_mut(groupid)
257    }
258
259    fn groups(&self) -> &[ParameterGroup<A, D>] {
260        self.group_manager.groups()
261    }
262
263    fn groups_mut(&mut self) -> &mut [ParameterGroup<A, D>] {
264        self.group_manager.groups_mut()
265    }
266
267    fn step_group(
268        &mut self,
269        groupid: usize,
270        gradients: &[Array<A, D>],
271    ) -> Result<Vec<Array<A, D>>> {
272        self.step = self.step.saturating_add(1);
273        let group_step = {
274            let counter = self.group_steps.entry(groupid).or_insert(0);
275            *counter = counter.saturating_add(1);
276            *counter
277        };
278        self.step_group_internal(groupid, group_step, gradients)
279    }
280
281    fn set_group_learning_rate(&mut self, groupid: usize, lr: A) -> Result<()> {
282        let group = self.group_manager.get_group_mut(groupid)?;
283        group.config.learning_rate = Some(lr);
284        Ok(())
285    }
286
287    fn set_group_weight_decay(&mut self, groupid: usize, wd: A) -> Result<()> {
288        let group = self.group_manager.get_group_mut(groupid)?;
289        group.config.weight_decay = Some(wd);
290        Ok(())
291    }
292}
293
294// Standard optimizer implementation for default behavior
295impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync> Optimizer<A, D>
296    for GroupedAdam<A, D>
297{
298    fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
299        if params.shape() != gradients.shape() {
300            return Err(OptimError::DimensionMismatch(format!(
301                "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
302                params.shape(),
303                gradients.shape()
304            )));
305        }
306
307        // Reuse a single implicit group across calls. Creating a new group on every
308        // call would reset the moment estimates, grow the group list without bound and
309        // make each step cost O(number of previous steps).
310        let reusable = self
311            .implicit_group
312            .filter(|id| self.group_manager.get_group(*id).is_ok());
313
314        let groupid = if let Some(id) = reusable {
315            let group = self.group_manager.get_group_mut(id)?;
316            let shape_changed =
317                group.params.len() != 1 || group.params[0].raw_dim() != params.raw_dim();
318            if shape_changed {
319                // A different parameter shape means the cached moments are meaningless.
320                group.params = vec![params.clone()];
321                group.state.clear();
322                self.group_steps.insert(id, 0);
323            } else {
324                group.params[0] = params.clone();
325            }
326            id
327        } else {
328            let id = self.add_group(vec![params.clone()], ParameterGroupConfig::new())?;
329            self.implicit_group = Some(id);
330            self.group_steps.insert(id, 0);
331            id
332        };
333
334        let result = self.step_group(groupid, std::slice::from_ref(gradients))?;
335
336        result.into_iter().next().ok_or_else(|| {
337            OptimError::InvalidConfig("grouped Adam step produced no parameters".to_string())
338        })
339    }
340
341    fn get_learning_rate(&self) -> A {
342        self.defaultlr
343    }
344
345    fn set_learning_rate(&mut self, learning_rate: A) {
346        self.defaultlr = learning_rate;
347    }
348}
349
350#[cfg(test)]
351mod tests {
352    use super::*;
353    use scirs2_core::ndarray::Array1;
354
355    #[test]
356    fn test_grouped_adam_creation() {
357        let optimizer: GroupedAdam<f64, scirs2_core::ndarray::Ix1> = GroupedAdam::new(0.001);
358        assert_eq!(optimizer.defaultlr, 0.001);
359        assert_eq!(optimizer.default_beta1, 0.9);
360        assert_eq!(optimizer.default_beta2, 0.999);
361    }
362
363    #[test]
364    fn test_grouped_adam_multiple_groups() {
365        let mut optimizer = GroupedAdam::new(0.001);
366
367        // Add first group with high learning rate
368        let params1 = vec![Array1::from_vec(vec![1.0, 2.0])];
369        let config1 = ParameterGroupConfig::new().with_learning_rate(0.01);
370        let group1 = optimizer
371            .add_group(params1, config1)
372            .expect("add_group succeeds in test_grouped_adam_multiple_groups");
373
374        // Add second group with low learning rate
375        let params2 = vec![Array1::from_vec(vec![3.0, 4.0, 5.0])];
376        let config2 = ParameterGroupConfig::new().with_learning_rate(0.0001);
377        let group2 = optimizer
378            .add_group(params2, config2)
379            .expect("add_group succeeds in test_grouped_adam_multiple_groups");
380
381        // Update first group
382        let grads1 = vec![Array1::from_vec(vec![0.1, 0.2])];
383        let updated1 = optimizer
384            .step_group(group1, &grads1)
385            .expect("step_group succeeds in test_grouped_adam_multiple_groups");
386
387        // Update second group
388        let grads2 = vec![Array1::from_vec(vec![0.3, 0.4, 0.5])];
389        let updated2 = optimizer
390            .step_group(group2, &grads2)
391            .expect("step_group succeeds in test_grouped_adam_multiple_groups");
392
393        // Verify different updates due to different learning rates
394        assert!(updated1[0][0] < 1.0); // Should decrease more
395        assert!(updated2[0][0] > 2.9); // Should decrease less
396    }
397
398    #[test]
399    fn test_grouped_adam_custom_betas() {
400        let mut optimizer = GroupedAdam::new(0.001);
401
402        // Add group with custom betas
403        let params = vec![Array1::from_vec(vec![1.0, 2.0])];
404        let config = ParameterGroupConfig::new()
405            .with_custom_param("beta1".to_string(), 0.8)
406            .with_custom_param("beta2".to_string(), 0.99);
407        let group = optimizer
408            .add_group(params, config)
409            .expect("optimizer.add_group succeeds in test_grouped_adam_custom_betas");
410
411        // Verify custom parameters are used
412        let group_ref = optimizer
413            .get_group(group)
414            .expect("optimizer.get_group succeeds in test_grouped_adam_custom_betas");
415        assert_eq!(group_ref.get_custom_param("beta1", 0.0), 0.8);
416        assert_eq!(group_ref.get_custom_param("beta2", 0.0), 0.99);
417    }
418
419    #[test]
420    fn test_grouped_adam_clear() {
421        let mut optimizer = GroupedAdam::new(0.001);
422
423        // Add groups
424        let params1 = vec![Array1::zeros(2)];
425        let config1 = ParameterGroupConfig::new();
426        optimizer
427            .add_group(params1, config1)
428            .expect("add_group succeeds in test_grouped_adam_clear");
429
430        assert_eq!(optimizer.groups().len(), 1);
431
432        // Clear groups
433        optimizer.clear_groups();
434
435        assert_eq!(optimizer.groups().len(), 0);
436        assert_eq!(optimizer.step, 0);
437    }
438
439    /// Regression test: `Optimizer::step` used to append a brand new parameter group on
440    /// every call, which reset the Adam moments, leaked memory without bound and made
441    /// each step cost O(previous steps).
442    #[test]
443    fn test_grouped_adam_step_reuses_single_group() {
444        let mut optimizer: GroupedAdam<f64, scirs2_core::ndarray::Ix1> = GroupedAdam::new(0.1);
445
446        let mut params = Array1::from_vec(vec![0.0f64]);
447        let gradients = Array1::from_vec(vec![1.0f64]);
448
449        for _ in 0..100 {
450            params = optimizer
451                .step(&params, &gradients)
452                .expect("step should succeed");
453        }
454
455        // Exactly one implicit group, not one per step.
456        assert_eq!(optimizer.groups().len(), 1);
457        assert_eq!(optimizer.group_step_count(0), 100);
458
459        // The moments must have been carried across the 100 steps.
460        let group = optimizer.get_group(0).expect("implicit group must exist");
461        let m_t = group
462            .state
463            .get("m_t")
464            .expect("first moment state must exist");
465        // m converges to the (constant) gradient value 1.0
466        assert!(
467            (m_t[0][0] - 1.0).abs() < 1e-3,
468            "first moment did not accumulate: {}",
469            m_t[0][0]
470        );
471    }
472
473    /// The first step of a group must use t = 1 (bias correction), not t = 2.
474    #[test]
475    fn test_grouped_adam_first_step_uses_t_one() {
476        let mut optimizer: GroupedAdam<f64, scirs2_core::ndarray::Ix1> = GroupedAdam::new(0.1);
477
478        let params = Array1::from_vec(vec![0.0f64]);
479        let gradients = Array1::from_vec(vec![1.0f64]);
480
481        let updated = optimizer
482            .step(&params, &gradients)
483            .expect("step should succeed");
484
485        // At t = 1 with a unit gradient Adam moves exactly -lr.
486        assert!(
487            (updated[0] + 0.1).abs() < 1e-9,
488            "expected -0.1, got {}",
489            updated[0]
490        );
491    }
492
493    /// Two explicit groups must keep independent bias-correction clocks.
494    #[test]
495    fn test_grouped_adam_per_group_step_counters() {
496        let mut optimizer = GroupedAdam::new(0.1);
497
498        let group_a = optimizer
499            .add_group(
500                vec![Array1::from_vec(vec![0.0f64])],
501                ParameterGroupConfig::new(),
502            )
503            .expect("add group a");
504        let group_b = optimizer
505            .add_group(
506                vec![Array1::from_vec(vec![0.0f64])],
507                ParameterGroupConfig::new(),
508            )
509            .expect("add group b");
510
511        let grads = vec![Array1::from_vec(vec![1.0f64])];
512
513        // Step group A five times before group B is touched at all.
514        for _ in 0..5 {
515            optimizer.step_group(group_a, &grads).expect("step group a");
516        }
517
518        let first_b = optimizer.step_group(group_b, &grads).expect("step group b");
519
520        assert_eq!(optimizer.group_step_count(group_a), 5);
521        assert_eq!(optimizer.group_step_count(group_b), 1);
522
523        // Group B is on its own first step, so it must move exactly -lr.
524        assert!(
525            (first_b[0][0] + 0.1).abs() < 1e-9,
526            "group B was contaminated by group A's clock: {}",
527            first_b[0][0]
528        );
529    }
530}