Skip to main content

optirs_core/optimizer_composition/
mod.rs

1// Optimizer composition framework
2//
3// This module provides compositions of optimizers to create more sophisticated
4// optimization strategies. It includes three main types of compositions:
5//
6// 1. **Sequential**: Apply multiple optimizers in sequence
7// 2. **Parallel**: Apply different optimizers to different parameter groups
8// 3. **Chained**: Wrap an optimizer with another (similar to Lookahead wrapping other optimizers)
9
10use crate::error::{OptimError, Result};
11use crate::optimizers::Optimizer;
12use scirs2_core::ndarray::{Array, Dimension, ScalarOperand};
13use scirs2_core::numeric::Float;
14use std::fmt::Debug;
15
16/// A sequential composition of optimizers
17///
18/// This applies multiple optimizers in sequence to the same parameters.
19/// Each optimizer's output becomes the input to the next optimizer.
20///
21/// # Example
22///
23/// ```
24/// use scirs2_core::ndarray::Array1;
25/// use optirs_core::optimizer_composition::SequentialOptimizer;
26/// use optirs_core::optimizers::{SGD, Adam, Optimizer};
27///
28/// // Create optimizers
29/// let sgd = SGD::new(0.1);
30/// let adam = Adam::new(0.01);
31///
32/// // Combine them sequentially
33/// let mut seq_optimizer = SequentialOptimizer::new(vec![
34///     Box::new(sgd),
35///     Box::new(adam),
36/// ]);
37///
38/// // Use the sequential optimizer
39/// let params = Array1::zeros(5);
40/// let gradients = Array1::ones(5);
41/// let updated_params = seq_optimizer.step(&params, &gradients).expect("seq_optimizer.step succeeds");
42/// ```
43pub struct SequentialOptimizer<A, D>
44where
45    A: Float + ScalarOperand + Debug,
46    D: Dimension,
47{
48    /// List of optimizers to apply in sequence
49    optimizers: Vec<Box<dyn Optimizer<A, D>>>,
50}
51
52impl<A, D> SequentialOptimizer<A, D>
53where
54    A: Float + ScalarOperand + Debug,
55    D: Dimension,
56{
57    /// Create a new sequential optimizer
58    ///
59    /// # Arguments
60    ///
61    /// * `optimizers` - List of optimizers to apply in sequence
62    pub fn new(optimizers: Vec<Box<dyn Optimizer<A, D>>>) -> Self {
63        Self { optimizers }
64    }
65
66    /// Add an optimizer to the sequence
67    ///
68    /// # Arguments
69    ///
70    /// * `optimizer` - The optimizer to add
71    pub fn add_optimizer(&mut self, optimizer: Box<dyn Optimizer<A, D>>) {
72        self.optimizers.push(optimizer);
73    }
74
75    /// Get the number of optimizers in the sequence
76    pub fn num_optimizers(&self) -> usize {
77        self.optimizers.len()
78    }
79
80    /// Get a reference to an optimizer by index
81    ///
82    /// # Arguments
83    ///
84    /// * `index` - The index of the optimizer
85    ///
86    /// # Returns
87    ///
88    /// A reference to the optimizer at the given index, or None if out of bounds
89    pub fn get_optimizer(&self, index: usize) -> Option<&dyn Optimizer<A, D>> {
90        if index < self.optimizers.len() {
91            Some(self.optimizers[index].as_ref())
92        } else {
93            None
94        }
95    }
96
97    /// Get a mutable reference to an optimizer by index
98    ///
99    /// # Arguments
100    ///
101    /// * `index` - The index of the optimizer
102    ///
103    /// # Returns
104    ///
105    /// A mutable reference to the optimizer at the given index, or None if out of bounds
106    pub fn get_optimizer_mut(&mut self, index: usize) -> Option<&mut dyn Optimizer<A, D>> {
107        if index < self.optimizers.len() {
108            Some(self.optimizers[index].as_mut())
109        } else {
110            None
111        }
112    }
113}
114
115impl<A, D> Optimizer<A, D> for SequentialOptimizer<A, D>
116where
117    A: Float + ScalarOperand + Debug,
118    D: Dimension,
119{
120    fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
121        // Check if we have any optimizers
122        if self.optimizers.is_empty() {
123            return Err(OptimError::InvalidConfig(
124                "SequentialOptimizer has no optimizers".to_string(),
125            ));
126        }
127
128        // Start with the initial parameters
129        let mut current_params = params.clone();
130
131        // Apply each optimizer in sequence
132        for optimizer in &mut self.optimizers {
133            current_params = optimizer.step(&current_params, gradients)?;
134        }
135
136        Ok(current_params)
137    }
138
139    fn get_learning_rate(&self) -> A {
140        // Return the learning rate of the first optimizer, or a default if empty
141        match self.optimizers.first() {
142            Some(optimizer) => optimizer.get_learning_rate(),
143            // Default learning rate; falls back to zero for exotic float types
144            None => A::from(0.01).unwrap_or_else(A::zero),
145        }
146    }
147
148    fn set_learning_rate(&mut self, learningrate: A) {
149        // Set the learning _rate for all optimizers
150        for optimizer in &mut self.optimizers {
151            optimizer.set_learning_rate(learningrate);
152        }
153    }
154}
155
156/// A struct for assigning parameters to specific groups for parallel optimization
157pub struct ParameterGroup<A, D>
158where
159    A: Float + ScalarOperand + Debug,
160    D: Dimension,
161{
162    /// The parameters in this group
163    pub params: Array<A, D>,
164    /// The index of the optimizer to use for this group
165    pub optimizerindex: usize,
166}
167
168impl<A, D> ParameterGroup<A, D>
169where
170    A: Float + ScalarOperand + Debug,
171    D: Dimension,
172{
173    /// Create a new parameter group
174    ///
175    /// # Arguments
176    ///
177    /// * `params` - The parameters in this group
178    /// * `optimizerindex` - The index of the optimizer to use for this group
179    pub fn new(params: Array<A, D>, optimizerindex: usize) -> Self {
180        Self {
181            params,
182            optimizerindex,
183        }
184    }
185}
186
187/// A parallel composition of optimizers
188///
189/// This applies different optimizers to different parameter groups.
190/// Each group of parameters is updated using its assigned optimizer.
191///
192/// # Example
193///
194/// ```
195/// use scirs2_core::ndarray::Array1;
196/// use optirs_core::optimizer_composition::{ParallelOptimizer, ParameterGroup};
197/// use optirs_core::optimizers::{SGD, Adam, Optimizer};
198///
199/// // Create optimizers
200/// let sgd = SGD::new(0.1);
201/// let adam = Adam::new(0.01);
202///
203/// // Create parameter groups
204/// let params1 = Array1::zeros(3);
205/// let params2 = Array1::zeros(5);
206///
207/// let group1 = ParameterGroup::new(params1, 0); // Use SGD
208/// let group2 = ParameterGroup::new(params2, 1); // Use Adam
209///
210/// // Combine them in parallel
211/// let mut parallel_optimizer = ParallelOptimizer::new(
212///     vec![Box::new(sgd), Box::new(adam)],
213///     vec![group1, group2],
214/// );
215///
216/// // The step method will update all parameter groups using their assigned optimizers
217/// // (In a real use case, you'd provide the corresponding gradients)
218/// ```
219pub struct ParallelOptimizer<A, D>
220where
221    A: Float + ScalarOperand + Debug,
222    D: Dimension,
223{
224    /// List of optimizers to apply to different parameter groups
225    optimizers: Vec<Box<dyn Optimizer<A, D>>>,
226    /// Groups of parameters with their assigned optimizer indices
227    parameter_groups: Vec<ParameterGroup<A, D>>,
228}
229
230impl<A, D> ParallelOptimizer<A, D>
231where
232    A: Float + ScalarOperand + Debug,
233    D: Dimension,
234{
235    /// Create a new parallel optimizer
236    ///
237    /// # Arguments
238    ///
239    /// * `optimizers` - List of optimizers to use
240    /// * `parameter_groups` - Groups of parameters with their assigned optimizer indices
241    pub fn new(
242        optimizers: Vec<Box<dyn Optimizer<A, D>>>,
243        parameter_groups: Vec<ParameterGroup<A, D>>,
244    ) -> Self {
245        Self {
246            optimizers,
247            parameter_groups,
248        }
249    }
250
251    /// Add an optimizer
252    ///
253    /// # Arguments
254    ///
255    /// * `optimizer` - The optimizer to add
256    ///
257    /// # Returns
258    ///
259    /// The index of the added optimizer
260    pub fn add_optimizer(&mut self, optimizer: Box<dyn Optimizer<A, D>>) -> usize {
261        let index = self.optimizers.len();
262        self.optimizers.push(optimizer);
263        index
264    }
265
266    /// Add a parameter group
267    ///
268    /// # Arguments
269    ///
270    /// * `params` - The parameters in this group
271    /// * `optimizerindex` - The index of the optimizer to use for this group
272    ///
273    /// # Returns
274    ///
275    /// Result with the index of the added parameter group, or an error if the optimizer index is invalid
276    pub fn add_parameter_group(
277        &mut self,
278        params: Array<A, D>,
279        optimizerindex: usize,
280    ) -> Result<usize> {
281        // Check if the optimizer _index is valid
282        if optimizerindex >= self.optimizers.len() {
283            return Err(OptimError::InvalidConfig(format!(
284                "Invalid optimizer _index: {}. Only {} optimizers available.",
285                optimizerindex,
286                self.optimizers.len()
287            )));
288        }
289
290        let _index = self.parameter_groups.len();
291        self.parameter_groups
292            .push(ParameterGroup::new(params, optimizerindex));
293        Ok(_index)
294    }
295
296    /// Get the number of optimizers
297    pub fn num_optimizers(&self) -> usize {
298        self.optimizers.len()
299    }
300
301    /// Get the number of parameter groups
302    pub fn num_parameter_groups(&self) -> usize {
303        self.parameter_groups.len()
304    }
305
306    /// Get a reference to an optimizer by index
307    ///
308    /// # Arguments
309    ///
310    /// * `index` - The index of the optimizer
311    ///
312    /// # Returns
313    ///
314    /// A reference to the optimizer at the given index, or None if out of bounds
315    pub fn get_optimizer(&self, index: usize) -> Option<&dyn Optimizer<A, D>> {
316        if index < self.optimizers.len() {
317            Some(self.optimizers[index].as_ref())
318        } else {
319            None
320        }
321    }
322
323    /// Get a mutable reference to an optimizer by index
324    ///
325    /// # Arguments
326    ///
327    /// * `index` - The index of the optimizer
328    ///
329    /// # Returns
330    ///
331    /// A mutable reference to the optimizer at the given index, or None if out of bounds
332    pub fn get_optimizer_mut(&mut self, index: usize) -> Option<&mut dyn Optimizer<A, D>> {
333        if index < self.optimizers.len() {
334            Some(self.optimizers[index].as_mut())
335        } else {
336            None
337        }
338    }
339
340    /// Get a reference to a parameter group by index
341    ///
342    /// # Arguments
343    ///
344    /// * `index` - The index of the parameter group
345    ///
346    /// # Returns
347    ///
348    /// A reference to the parameter group at the given index, or None if out of bounds
349    pub fn get_parameter_group(&self, index: usize) -> Option<&ParameterGroup<A, D>> {
350        self.parameter_groups.get(index)
351    }
352
353    /// Get a mutable reference to a parameter group by index
354    ///
355    /// # Arguments
356    ///
357    /// * `index` - The index of the parameter group
358    ///
359    /// # Returns
360    ///
361    /// A mutable reference to the parameter group at the given index, or None if out of bounds
362    pub fn get_parameter_group_mut(&mut self, index: usize) -> Option<&mut ParameterGroup<A, D>> {
363        self.parameter_groups.get_mut(index)
364    }
365
366    /// Get all current parameter values as a single array
367    ///
368    /// # Returns
369    ///
370    /// A result containing all parameter values concatenated into a single array
371    pub fn get_all_parameters(&self) -> Result<Vec<Array<A, D>>> {
372        Ok(self
373            .parameter_groups
374            .iter()
375            .map(|group| group.params.clone())
376            .collect())
377    }
378
379    /// Update all parameter groups using their assigned optimizers
380    ///
381    /// # Arguments
382    ///
383    /// * `gradients` - List of gradient arrays corresponding to parameter groups
384    ///
385    /// # Returns
386    ///
387    /// Result with the updated parameter values, or an error
388    pub fn update_all_parameters(&mut self, gradients: &[Array<A, D>]) -> Result<Vec<Array<A, D>>> {
389        // Check if the number of gradients matches the number of parameter groups
390        if gradients.len() != self.parameter_groups.len() {
391            return Err(OptimError::InvalidConfig(format!(
392                "Number of gradients ({}) does not match number of parameter groups ({})",
393                gradients.len(),
394                self.parameter_groups.len()
395            )));
396        }
397
398        let mut updated_params = Vec::with_capacity(self.parameter_groups.len());
399
400        // Update each parameter group using its assigned optimizer
401        for (i, group) in self.parameter_groups.iter_mut().enumerate() {
402            let optimizerindex = group.optimizerindex;
403
404            // Check if the optimizer index is valid
405            if optimizerindex >= self.optimizers.len() {
406                return Err(OptimError::InvalidConfig(format!(
407                    "Invalid optimizer index: {}. Only {} optimizers available.",
408                    optimizerindex,
409                    self.optimizers.len()
410                )));
411            }
412
413            // Get the optimizer and update the parameters
414            let optimizer = &mut self.optimizers[optimizerindex];
415            let params = &group.params;
416            let gradient = &gradients[i];
417
418            // Update the parameters
419            let updated = optimizer.step(params, gradient)?;
420            group.params = updated.clone();
421            updated_params.push(updated);
422        }
423
424        Ok(updated_params)
425    }
426}
427
428impl<A, D> Optimizer<A, D> for ParallelOptimizer<A, D>
429where
430    A: Float + ScalarOperand + Debug,
431    D: Dimension,
432{
433    fn step(&mut self, _params: &Array<A, D>, _gradients: &Array<A, D>) -> Result<Array<A, D>> {
434        // This implementation is a bit tricky since we have multiple parameter groups
435        // We'll return an error message directing users to use update_all_parameters instead
436        Err(OptimError::InvalidConfig(
437            "ParallelOptimizer doesn't support the standard step method. Use update_all_parameters instead."
438                .to_string(),
439        ))
440    }
441
442    /// Updates several parameter tensors, one per parameter group
443    ///
444    /// Existing parameter groups are **reused**: only their parameter values are
445    /// refreshed, so the optimizer assignment made by
446    /// [`ParallelOptimizer::add_parameter_group`] survives across calls. Groups are
447    /// rebuilt only when the caller changes the number of tensors or their shapes, in
448    /// which case tensor `i` is assigned to optimizer `min(i, optimizers.len() - 1)`.
449    fn step_list(
450        &mut self,
451        params_list: &[&Array<A, D>],
452        gradients_list: &[&Array<A, D>],
453    ) -> Result<Vec<Array<A, D>>> {
454        if params_list.len() != gradients_list.len() {
455            return Err(OptimError::InvalidConfig(format!(
456                "Number of parameter arrays ({}) does not match number of gradient arrays ({})",
457                params_list.len(),
458                gradients_list.len()
459            )));
460        }
461
462        // Guard against an empty optimizer list: the fallback assignment below would
463        // otherwise underflow when computing `optimizers.len() - 1`.
464        let last_optimizer = self.optimizers.len().checked_sub(1).ok_or_else(|| {
465            OptimError::InvalidConfig(
466                "ParallelOptimizer has no optimizers; add at least one with add_optimizer \
467                 before calling step_list."
468                    .to_string(),
469            )
470        })?;
471
472        // Reuse the existing groups when they still describe the same tensors, so that
473        // per-group optimizer assignments are not silently discarded on every call.
474        let layout_matches = self.parameter_groups.len() == params_list.len()
475            && self
476                .parameter_groups
477                .iter()
478                .zip(params_list.iter())
479                .all(|(group, params)| group.params.raw_dim() == params.raw_dim());
480
481        if layout_matches {
482            for (group, params) in self.parameter_groups.iter_mut().zip(params_list.iter()) {
483                group.params = (*params).clone();
484            }
485        } else {
486            self.parameter_groups = params_list
487                .iter()
488                .enumerate()
489                .map(|(i, params)| {
490                    // Use the last optimizer for any tensor beyond the optimizer list
491                    ParameterGroup::new((*params).clone(), i.min(last_optimizer))
492                })
493                .collect();
494        }
495
496        // Convert gradients_list to owned arrays
497        let gradients_vec: Vec<Array<A, D>> = gradients_list.iter().map(|&g| g.clone()).collect();
498
499        // Update parameter groups using their assigned optimizers
500        self.update_all_parameters(&gradients_vec)
501    }
502
503    fn get_learning_rate(&self) -> A {
504        // Return the learning rate of the first optimizer, or a default if empty
505        if let Some(optimizer) = self.optimizers.first() {
506            optimizer.get_learning_rate()
507        } else {
508            // Default learning rate: 0.01 always fits in A (f32/f64)
509            A::from(0.01).expect("SequentialOptimizer: default learning rate (0.01) must fit in A")
510        }
511    }
512
513    fn set_learning_rate(&mut self, learningrate: A) {
514        // Set the learning _rate for all optimizers
515        for optimizer in &mut self.optimizers {
516            optimizer.set_learning_rate(learningrate);
517        }
518    }
519}
520
521/// A chained composition of optimizers
522///
523/// This wraps one optimizer with another, similar to how Lookahead wraps
524/// another optimizer. The inner optimizer is applied first, and then the
525/// outer optimizer is applied to the result.
526///
527/// # Example
528///
529/// ```
530/// use scirs2_core::ndarray::Array1;
531/// use optirs_core::optimizer_composition::ChainedOptimizer;
532/// use optirs_core::optimizers::{SGD, Adam, Optimizer};
533///
534/// // Create optimizers
535/// let inner = SGD::new(0.1);
536/// let outer = Adam::new(0.01);
537///
538/// // Chain them together
539/// let mut chained_optimizer = ChainedOptimizer::new(Box::new(inner), Box::new(outer));
540///
541/// // Use the chained optimizer
542/// let params = Array1::zeros(5);
543/// let gradients = Array1::ones(5);
544/// let updated_params = chained_optimizer.step(&params, &gradients).expect("chained_optimizer.step succeeds");
545/// ```
546pub struct ChainedOptimizer<A, D>
547where
548    A: Float + ScalarOperand + Debug,
549    D: Dimension,
550{
551    /// The inner optimizer, applied first
552    inner: Box<dyn Optimizer<A, D>>,
553    /// The outer optimizer, applied to the result of the inner optimizer
554    outer: Box<dyn Optimizer<A, D>>,
555}
556
557impl<A, D> ChainedOptimizer<A, D>
558where
559    A: Float + ScalarOperand + Debug,
560    D: Dimension,
561{
562    /// Create a new chained optimizer
563    ///
564    /// # Arguments
565    ///
566    /// * `inner` - The inner optimizer, applied first
567    /// * `outer` - The outer optimizer, applied to the result of the inner optimizer
568    pub fn new(inner: Box<dyn Optimizer<A, D>>, outer: Box<dyn Optimizer<A, D>>) -> Self {
569        Self { inner, outer }
570    }
571
572    /// Get a reference to the inner optimizer
573    pub fn inner(&self) -> &dyn Optimizer<A, D> {
574        self.inner.as_ref()
575    }
576
577    /// Get a mutable reference to the inner optimizer
578    pub fn inner_mut(&mut self) -> &mut dyn Optimizer<A, D> {
579        self.inner.as_mut()
580    }
581
582    /// Get a reference to the outer optimizer
583    pub fn outer(&self) -> &dyn Optimizer<A, D> {
584        self.outer.as_ref()
585    }
586
587    /// Get a mutable reference to the outer optimizer
588    pub fn outer_mut(&mut self) -> &mut dyn Optimizer<A, D> {
589        self.outer.as_mut()
590    }
591}
592
593impl<A, D> Optimizer<A, D> for ChainedOptimizer<A, D>
594where
595    A: Float + ScalarOperand + Debug,
596    D: Dimension,
597{
598    fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
599        // Apply the inner optimizer first
600        let intermediate_params = self.inner.step(params, gradients)?;
601
602        // Then apply the outer optimizer to the result
603        self.outer.step(&intermediate_params, gradients)
604    }
605
606    fn get_learning_rate(&self) -> A {
607        // Return the learning rate of the inner optimizer
608        self.inner.get_learning_rate()
609    }
610
611    fn set_learning_rate(&mut self, learningrate: A) {
612        // Set the learning _rate for both optimizers
613        self.inner.set_learning_rate(learningrate);
614        self.outer.set_learning_rate(learningrate);
615    }
616}
617
618/// A weighted composition of optimizers
619///
620/// Runs all optimizers on the same parameters/gradients and returns the
621/// weighted average of their outputs. This allows blending the behavior
622/// of multiple optimization strategies.
623///
624/// # Example
625///
626/// ```
627/// use scirs2_core::ndarray::Array1;
628/// use optirs_core::optimizer_composition::WeightedOptimizer;
629/// use optirs_core::optimizers::{SGD, Adam, Optimizer};
630///
631/// // Create a weighted combination of SGD and Adam
632/// let mut weighted = WeightedOptimizer::new()
633///     .add_optimizer(Box::new(SGD::new(0.1)), 0.7)
634///     .add_optimizer(Box::new(Adam::new(0.01)), 0.3);
635///
636/// let params = Array1::zeros(3);
637/// let gradients = Array1::ones(3);
638/// let updated = weighted.step(&params, &gradients).expect("step failed");
639/// ```
640pub struct WeightedOptimizer<A, D>
641where
642    A: Float + ScalarOperand + Debug,
643    D: Dimension,
644{
645    /// The optimizers and their associated weights
646    optimizers: Vec<Box<dyn Optimizer<A, D>>>,
647    /// The weight for each optimizer
648    weights: Vec<A>,
649}
650
651impl<A, D> Default for WeightedOptimizer<A, D>
652where
653    A: Float + ScalarOperand + Debug,
654    D: Dimension,
655{
656    fn default() -> Self {
657        Self::new()
658    }
659}
660
661impl<A, D> WeightedOptimizer<A, D>
662where
663    A: Float + ScalarOperand + Debug,
664    D: Dimension,
665{
666    /// Create a new empty weighted optimizer
667    pub fn new() -> Self {
668        Self {
669            optimizers: Vec::new(),
670            weights: Vec::new(),
671        }
672    }
673
674    /// Add an optimizer with a given weight (builder pattern)
675    ///
676    /// # Arguments
677    ///
678    /// * `opt` - The optimizer to add
679    /// * `weight` - The weight for this optimizer
680    pub fn add_optimizer(mut self, opt: Box<dyn Optimizer<A, D>>, weight: A) -> Self {
681        self.optimizers.push(opt);
682        self.weights.push(weight);
683        self
684    }
685
686    /// Add multiple optimizers at once (builder pattern)
687    ///
688    /// # Arguments
689    ///
690    /// * `opts` - A vector of (optimizer, weight) pairs
691    pub fn with_optimizers(mut self, opts: Vec<(Box<dyn Optimizer<A, D>>, A)>) -> Self {
692        for (opt, weight) in opts {
693            self.optimizers.push(opt);
694            self.weights.push(weight);
695        }
696        self
697    }
698
699    /// Normalize weights so they sum to 1
700    pub fn normalize_weights(&mut self) {
701        let sum: A = self.weights.iter().copied().fold(A::zero(), |a, b| a + b);
702        if sum > A::zero() {
703            for w in &mut self.weights {
704                *w = *w / sum;
705            }
706        }
707    }
708
709    /// Get the number of optimizers
710    pub fn num_optimizers(&self) -> usize {
711        self.optimizers.len()
712    }
713
714    /// Get the current weights
715    pub fn weights(&self) -> &[A] {
716        &self.weights
717    }
718}
719
720impl<A, D> Optimizer<A, D> for WeightedOptimizer<A, D>
721where
722    A: Float + ScalarOperand + Debug,
723    D: Dimension,
724{
725    fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
726        if self.optimizers.is_empty() {
727            return Err(OptimError::InvalidConfig(
728                "WeightedOptimizer has no optimizers".to_string(),
729            ));
730        }
731
732        // Compute the weight sum for normalization
733        let weight_sum: A = self.weights.iter().copied().fold(A::zero(), |a, b| a + b);
734        if weight_sum <= A::zero() {
735            return Err(OptimError::InvalidConfig(
736                "WeightedOptimizer weight sum must be positive".to_string(),
737            ));
738        }
739
740        // Run each optimizer and accumulate the weighted result
741        let mut result: Option<Array<A, D>> = None;
742
743        for (optimizer, &weight) in self.optimizers.iter_mut().zip(self.weights.iter()) {
744            let updated = optimizer.step(params, gradients)?;
745            let normalized_weight = weight / weight_sum;
746
747            match result {
748                None => {
749                    result = Some(updated * normalized_weight);
750                }
751                Some(ref mut acc) => {
752                    acc.zip_mut_with(&updated, |a, &b| {
753                        *a = *a + b * normalized_weight;
754                    });
755                }
756            }
757        }
758
759        result.ok_or_else(|| {
760            OptimError::InvalidConfig("WeightedOptimizer produced no result".to_string())
761        })
762    }
763
764    fn get_learning_rate(&self) -> A {
765        if let Some(optimizer) = self.optimizers.first() {
766            optimizer.get_learning_rate()
767        } else {
768            A::from(0.01).expect("failed to convert default learning rate")
769        }
770    }
771
772    fn set_learning_rate(&mut self, learning_rate: A) {
773        for optimizer in &mut self.optimizers {
774            optimizer.set_learning_rate(learning_rate);
775        }
776    }
777}
778
779#[cfg(test)]
780mod tests {
781    use super::*;
782    use crate::optimizers::{Adam, SGD};
783    use approx::assert_abs_diff_eq;
784    use scirs2_core::ndarray::Array1;
785
786    #[test]
787    fn test_sequential_optimizer() {
788        // Create a sequential optimizer with SGD followed by Adam
789        let sgd = SGD::new(0.1);
790        let adam = Adam::new(0.01);
791
792        let mut seq_optimizer: SequentialOptimizer<f64, scirs2_core::ndarray::Ix1> =
793            SequentialOptimizer::new(vec![Box::new(sgd), Box::new(adam)]);
794
795        // Create test parameters and gradients
796        let params = Array1::zeros(3);
797        let gradients = Array1::from_vec(vec![1.0, 2.0, 3.0]);
798
799        // Apply the sequential optimizer
800        let updated_params = seq_optimizer
801            .step(&params, &gradients)
802            .expect("step succeeds in test_sequential_optimizer");
803
804        // Verify the result
805        // First SGD updates: params - 0.1 * gradients = [0, 0, 0] - 0.1 * [1, 2, 3] = [-0.1, -0.2, -0.3]
806        // Then Adam makes additional updates
807        assert!(updated_params[0] < -0.1);
808        assert!(updated_params[1] < -0.2);
809        assert!(updated_params[2] < -0.3);
810    }
811
812    #[test]
813    fn test_parallel_optimizer() {
814        // Create a parallel optimizer with SGD and Adam
815        let sgd = SGD::new(0.1);
816        let adam = Adam::new(0.01);
817
818        let params1 = Array1::zeros(2);
819        let params2 = Array1::zeros(3);
820
821        let group1 = ParameterGroup::new(params1.clone(), 0); // Use SGD
822        let group2 = ParameterGroup::new(params2.clone(), 1); // Use Adam
823
824        let mut parallel_optimizer: ParallelOptimizer<f64, scirs2_core::ndarray::Ix1> =
825            ParallelOptimizer::new(vec![Box::new(sgd), Box::new(adam)], vec![group1, group2]);
826
827        // Create test gradients
828        let gradients1 = Array1::from_vec(vec![1.0, 2.0]);
829        let gradients2 = Array1::from_vec(vec![3.0, 4.0, 5.0]);
830
831        // Update the parameters
832        let updated_params = parallel_optimizer
833            .update_all_parameters(&[gradients1, gradients2])
834            .expect("update_all_parameters succeeds in test_parallel_optimizer");
835
836        // Verify the results
837        // Group 1 (SGD): params - 0.1 * gradients = [0, 0] - 0.1 * [1, 2] = [-0.1, -0.2]
838        assert_abs_diff_eq!(updated_params[0][0], -0.1);
839        assert_abs_diff_eq!(updated_params[0][1], -0.2);
840
841        // Group 2 (Adam): The update will be different due to Adam's adaptive nature
842        // Just verify it's different from the original params
843        assert!(updated_params[1][0] != 0.0);
844        assert!(updated_params[1][1] != 0.0);
845        assert!(updated_params[1][2] != 0.0);
846    }
847
848    #[test]
849    fn test_chained_optimizer() {
850        // Create a chained optimizer with SGD as inner and Adam as outer
851        let inner = SGD::new(0.1);
852        let outer = Adam::new(0.01);
853
854        let mut chained_optimizer: ChainedOptimizer<f64, scirs2_core::ndarray::Ix1> =
855            ChainedOptimizer::new(Box::new(inner), Box::new(outer));
856
857        // Create test parameters and gradients
858        let params = Array1::zeros(3);
859        let gradients = Array1::from_vec(vec![1.0, 2.0, 3.0]);
860
861        // Apply the chained optimizer
862        let updated_params = chained_optimizer
863            .step(&params, &gradients)
864            .expect("step succeeds in test_chained_optimizer");
865
866        // Verify the result
867        // Inner (SGD): params - 0.1 * gradients = [0, 0, 0] - 0.1 * [1, 2, 3] = [-0.1, -0.2, -0.3]
868        // Then outer (Adam) applies another update
869        assert!(updated_params[0] < -0.1);
870        assert!(updated_params[1] < -0.2);
871        assert!(updated_params[2] < -0.3);
872    }
873
874    #[test]
875    fn test_sequential_learning_rate() {
876        // Create a sequential optimizer with SGD followed by Adam
877        let sgd = SGD::new(0.1);
878        let adam = Adam::new(0.01);
879
880        let mut seq_optimizer: SequentialOptimizer<f64, scirs2_core::ndarray::Ix1> =
881            SequentialOptimizer::new(vec![Box::new(sgd), Box::new(adam)]);
882
883        // Test getting the learning rate (should be from the first optimizer)
884        assert_abs_diff_eq!(seq_optimizer.get_learning_rate(), 0.1);
885
886        // Test setting the learning rate for all optimizers
887        seq_optimizer.set_learning_rate(0.05);
888
889        // Verify the learning rate has been set for both optimizers
890        assert_abs_diff_eq!(seq_optimizer.get_learning_rate(), 0.05);
891        assert_abs_diff_eq!(
892            seq_optimizer
893                .get_optimizer(0)
894                .expect("get_optimizer succeeds in test_sequential_learning_rate")
895                .get_learning_rate(),
896            0.05
897        );
898        assert_abs_diff_eq!(
899            seq_optimizer
900                .get_optimizer(1)
901                .expect("get_optimizer succeeds in test_sequential_learning_rate")
902                .get_learning_rate(),
903            0.05
904        );
905    }
906
907    #[test]
908    fn test_parallel_optimizer_step_list() {
909        // Create a parallel optimizer with SGD and Adam
910        let sgd = SGD::new(0.1);
911        let adam = Adam::new(0.01);
912
913        let mut parallel_optimizer: ParallelOptimizer<f64, scirs2_core::ndarray::Ix1> =
914            ParallelOptimizer::new(vec![Box::new(sgd), Box::new(adam)], vec![]);
915
916        // Create test parameters and gradients
917        let params1 = Array1::zeros(2);
918        let params2 = Array1::zeros(3);
919        let params3 = Array1::zeros(4);
920
921        let gradients1 = Array1::from_vec(vec![1.0, 2.0]);
922        let gradients2 = Array1::from_vec(vec![3.0, 4.0, 5.0]);
923        let gradients3 = Array1::from_vec(vec![6.0, 7.0, 8.0, 9.0]);
924
925        // Use step_list to update all parameters
926        let params_refs = vec![&params1, &params2, &params3];
927        let gradients_refs = vec![&gradients1, &gradients2, &gradients3];
928
929        let updated_params = parallel_optimizer
930            .step_list(&params_refs, &gradients_refs)
931            .expect("step_list succeeds in test_parallel_optimizer_step_list");
932
933        // Verify the results
934        // Group 1 (SGD): params - 0.1 * gradients = [0, 0] - 0.1 * [1, 2] = [-0.1, -0.2]
935        assert_abs_diff_eq!(updated_params[0][0], -0.1);
936        assert_abs_diff_eq!(updated_params[0][1], -0.2);
937
938        // Group 2 will use SGD since we only have 2 optimizers and index 1 % 2 = 1 (Adam)
939        // Adam: The update will be different than SGD
940        assert!(updated_params[1][0] != -0.3);
941
942        // Group 3 will wrap around to optimize with Adam
943        // Just check that it's been updated from zero
944        assert!(updated_params[2][0] < 0.0);
945    }
946
947    #[test]
948    fn test_chained_optimizer_learning_rate() {
949        // Create a chained optimizer with SGD as inner and Adam as outer
950        let inner = SGD::new(0.1);
951        let outer = Adam::new(0.01);
952
953        let mut chained_optimizer: ChainedOptimizer<f64, scirs2_core::ndarray::Ix1> =
954            ChainedOptimizer::new(Box::new(inner), Box::new(outer));
955
956        // Test getting the learning rate (should be from the inner optimizer)
957        assert_abs_diff_eq!(chained_optimizer.get_learning_rate(), 0.1);
958
959        // Test setting the learning rate for both optimizers
960        chained_optimizer.set_learning_rate(0.05);
961
962        // Verify the learning rate has been set for both optimizers
963        assert_abs_diff_eq!(chained_optimizer.get_learning_rate(), 0.05);
964        assert_abs_diff_eq!(chained_optimizer.inner().get_learning_rate(), 0.05);
965        assert_abs_diff_eq!(chained_optimizer.outer().get_learning_rate(), 0.05);
966    }
967
968    #[test]
969    fn test_weighted_optimizer_basic() {
970        // Create two SGD optimizers with different learning rates
971        let sgd1 = SGD::new(0.1);
972        let sgd2 = SGD::new(0.2);
973
974        let mut weighted: WeightedOptimizer<f64, scirs2_core::ndarray::Ix1> =
975            WeightedOptimizer::new()
976                .add_optimizer(Box::new(sgd1), 0.5)
977                .add_optimizer(Box::new(sgd2), 0.5);
978
979        let params = Array1::zeros(3);
980        let gradients = Array1::ones(3);
981
982        let updated = weighted.step(&params, &gradients).expect("step failed");
983
984        // SGD1: params - 0.1 * grads = [-0.1, -0.1, -0.1]
985        // SGD2: params - 0.2 * grads = [-0.2, -0.2, -0.2]
986        // Weighted avg (0.5 each): 0.5*(-0.1) + 0.5*(-0.2) = -0.15
987        assert_abs_diff_eq!(updated[0], -0.15, epsilon = 1e-10);
988        assert_abs_diff_eq!(updated[1], -0.15, epsilon = 1e-10);
989        assert_abs_diff_eq!(updated[2], -0.15, epsilon = 1e-10);
990    }
991
992    #[test]
993    fn test_weighted_optimizer_unequal_weights() {
994        let sgd1 = SGD::new(0.1);
995        let sgd2 = SGD::new(0.2);
996
997        let mut weighted: WeightedOptimizer<f64, scirs2_core::ndarray::Ix1> =
998            WeightedOptimizer::new()
999                .add_optimizer(Box::new(sgd1), 3.0)
1000                .add_optimizer(Box::new(sgd2), 1.0);
1001
1002        let params = Array1::zeros(2);
1003        let gradients = Array1::ones(2);
1004
1005        let updated = weighted.step(&params, &gradients).expect("step failed");
1006
1007        // SGD1: [-0.1, -0.1], SGD2: [-0.2, -0.2]
1008        // Weights normalized: 3/4=0.75, 1/4=0.25
1009        // Result: 0.75*(-0.1) + 0.25*(-0.2) = -0.075 - 0.05 = -0.125
1010        assert_abs_diff_eq!(updated[0], -0.125, epsilon = 1e-10);
1011    }
1012
1013    #[test]
1014    fn test_weighted_optimizer_empty() {
1015        let mut weighted: WeightedOptimizer<f64, scirs2_core::ndarray::Ix1> =
1016            WeightedOptimizer::new();
1017
1018        let params = Array1::zeros(3);
1019        let gradients = Array1::ones(3);
1020
1021        let result = weighted.step(&params, &gradients);
1022        assert!(result.is_err());
1023    }
1024
1025    #[test]
1026    fn test_weighted_optimizer_normalize_weights() {
1027        let mut weighted: WeightedOptimizer<f64, scirs2_core::ndarray::Ix1> =
1028            WeightedOptimizer::new()
1029                .add_optimizer(Box::new(SGD::new(0.1)), 2.0)
1030                .add_optimizer(Box::new(SGD::new(0.2)), 8.0);
1031
1032        weighted.normalize_weights();
1033
1034        assert_abs_diff_eq!(weighted.weights()[0], 0.2, epsilon = 1e-10);
1035        assert_abs_diff_eq!(weighted.weights()[1], 0.8, epsilon = 1e-10);
1036    }
1037
1038    #[test]
1039    fn test_weighted_optimizer_learning_rate() {
1040        let mut weighted: WeightedOptimizer<f64, scirs2_core::ndarray::Ix1> =
1041            WeightedOptimizer::new()
1042                .add_optimizer(Box::new(SGD::new(0.1)), 1.0)
1043                .add_optimizer(Box::new(Adam::new(0.01)), 1.0);
1044
1045        // Learning rate comes from the first optimizer
1046        assert_abs_diff_eq!(weighted.get_learning_rate(), 0.1);
1047
1048        // Setting learning rate applies to all
1049        weighted.set_learning_rate(0.05);
1050        assert_abs_diff_eq!(weighted.get_learning_rate(), 0.05);
1051    }
1052
1053    #[test]
1054    fn test_weighted_optimizer_with_optimizers() {
1055        let opts: Vec<(Box<dyn Optimizer<f64, scirs2_core::ndarray::Ix1>>, f64)> = vec![
1056            (Box::new(SGD::new(0.1)), 1.0),
1057            (Box::new(SGD::new(0.2)), 1.0),
1058        ];
1059
1060        let weighted: WeightedOptimizer<f64, scirs2_core::ndarray::Ix1> =
1061            WeightedOptimizer::new().with_optimizers(opts);
1062
1063        assert_eq!(weighted.num_optimizers(), 2);
1064        assert_abs_diff_eq!(weighted.weights()[0], 1.0);
1065        assert_abs_diff_eq!(weighted.weights()[1], 1.0);
1066    }
1067
1068    /// Regression test for F88: `step_list` used to rebuild every parameter group on
1069    /// each call, discarding the optimizer assignment set up by
1070    /// `add_parameter_group`, and it underflowed on an empty optimizer list.
1071    #[test]
1072    fn test_parallel_optimizer_step_list_preserves_group_assignment() {
1073        let sgd = SGD::new(0.1);
1074        let adam = Adam::new(0.01);
1075
1076        let mut parallel_optimizer: ParallelOptimizer<f64, scirs2_core::ndarray::Ix1> =
1077            ParallelOptimizer::new(vec![Box::new(sgd), Box::new(adam)], vec![]);
1078
1079        // Deliberately assign BOTH tensors to optimizer 1 (Adam).
1080        parallel_optimizer
1081            .add_parameter_group(Array1::zeros(2), 1)
1082            .expect("add group 0");
1083        parallel_optimizer
1084            .add_parameter_group(Array1::zeros(2), 1)
1085            .expect("add group 1");
1086
1087        let params1 = Array1::zeros(2);
1088        let params2 = Array1::zeros(2);
1089        let grads1 = Array1::from_vec(vec![1.0, 2.0]);
1090        let grads2 = Array1::from_vec(vec![1.0, 2.0]);
1091
1092        let updated = parallel_optimizer
1093            .step_list(&[&params1, &params2], &[&grads1, &grads2])
1094            .expect("step_list failed");
1095
1096        // Both groups must still be assigned to Adam, so neither may show the plain
1097        // SGD result of -0.1 that the rebuilt-groups bug produced for group 0.
1098        assert_eq!(
1099            parallel_optimizer
1100                .get_parameter_group(0)
1101                .expect("group 0")
1102                .optimizerindex,
1103            1
1104        );
1105        assert_eq!(
1106            parallel_optimizer
1107                .get_parameter_group(1)
1108                .expect("group 1")
1109                .optimizerindex,
1110            1
1111        );
1112        assert!(
1113            (updated[0][0] + 0.1).abs() > 1e-6,
1114            "group 0 was silently reassigned to SGD: {}",
1115            updated[0][0]
1116        );
1117    }
1118
1119    /// An empty optimizer list must produce a clear error instead of underflowing.
1120    #[test]
1121    fn test_parallel_optimizer_step_list_rejects_empty_optimizers() {
1122        let mut parallel_optimizer: ParallelOptimizer<f64, scirs2_core::ndarray::Ix1> =
1123            ParallelOptimizer::new(vec![], vec![]);
1124
1125        let params = Array1::zeros(2);
1126        let grads = Array1::from_vec(vec![1.0, 2.0]);
1127
1128        let result = parallel_optimizer.step_list(&[&params], &[&grads]);
1129        assert!(result.is_err(), "empty optimizer list must be rejected");
1130    }
1131}