1use 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
16pub struct SequentialOptimizer<A, D>
44where
45 A: Float + ScalarOperand + Debug,
46 D: Dimension,
47{
48 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 pub fn new(optimizers: Vec<Box<dyn Optimizer<A, D>>>) -> Self {
63 Self { optimizers }
64 }
65
66 pub fn add_optimizer(&mut self, optimizer: Box<dyn Optimizer<A, D>>) {
72 self.optimizers.push(optimizer);
73 }
74
75 pub fn num_optimizers(&self) -> usize {
77 self.optimizers.len()
78 }
79
80 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 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 if self.optimizers.is_empty() {
123 return Err(OptimError::InvalidConfig(
124 "SequentialOptimizer has no optimizers".to_string(),
125 ));
126 }
127
128 let mut current_params = params.clone();
130
131 for optimizer in &mut self.optimizers {
133 current_params = optimizer.step(¤t_params, gradients)?;
134 }
135
136 Ok(current_params)
137 }
138
139 fn get_learning_rate(&self) -> A {
140 match self.optimizers.first() {
142 Some(optimizer) => optimizer.get_learning_rate(),
143 None => A::from(0.01).unwrap_or_else(A::zero),
145 }
146 }
147
148 fn set_learning_rate(&mut self, learningrate: A) {
149 for optimizer in &mut self.optimizers {
151 optimizer.set_learning_rate(learningrate);
152 }
153 }
154}
155
156pub struct ParameterGroup<A, D>
158where
159 A: Float + ScalarOperand + Debug,
160 D: Dimension,
161{
162 pub params: Array<A, D>,
164 pub optimizerindex: usize,
166}
167
168impl<A, D> ParameterGroup<A, D>
169where
170 A: Float + ScalarOperand + Debug,
171 D: Dimension,
172{
173 pub fn new(params: Array<A, D>, optimizerindex: usize) -> Self {
180 Self {
181 params,
182 optimizerindex,
183 }
184 }
185}
186
187pub struct ParallelOptimizer<A, D>
220where
221 A: Float + ScalarOperand + Debug,
222 D: Dimension,
223{
224 optimizers: Vec<Box<dyn Optimizer<A, D>>>,
226 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 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 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 pub fn add_parameter_group(
277 &mut self,
278 params: Array<A, D>,
279 optimizerindex: usize,
280 ) -> Result<usize> {
281 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 pub fn num_optimizers(&self) -> usize {
298 self.optimizers.len()
299 }
300
301 pub fn num_parameter_groups(&self) -> usize {
303 self.parameter_groups.len()
304 }
305
306 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 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 pub fn get_parameter_group(&self, index: usize) -> Option<&ParameterGroup<A, D>> {
350 self.parameter_groups.get(index)
351 }
352
353 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 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 pub fn update_all_parameters(&mut self, gradients: &[Array<A, D>]) -> Result<Vec<Array<A, D>>> {
389 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 for (i, group) in self.parameter_groups.iter_mut().enumerate() {
402 let optimizerindex = group.optimizerindex;
403
404 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 let optimizer = &mut self.optimizers[optimizerindex];
415 let params = &group.params;
416 let gradient = &gradients[i];
417
418 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 Err(OptimError::InvalidConfig(
437 "ParallelOptimizer doesn't support the standard step method. Use update_all_parameters instead."
438 .to_string(),
439 ))
440 }
441
442 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 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 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 ParameterGroup::new((*params).clone(), i.min(last_optimizer))
492 })
493 .collect();
494 }
495
496 let gradients_vec: Vec<Array<A, D>> = gradients_list.iter().map(|&g| g.clone()).collect();
498
499 self.update_all_parameters(&gradients_vec)
501 }
502
503 fn get_learning_rate(&self) -> A {
504 if let Some(optimizer) = self.optimizers.first() {
506 optimizer.get_learning_rate()
507 } else {
508 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 for optimizer in &mut self.optimizers {
516 optimizer.set_learning_rate(learningrate);
517 }
518 }
519}
520
521pub struct ChainedOptimizer<A, D>
547where
548 A: Float + ScalarOperand + Debug,
549 D: Dimension,
550{
551 inner: Box<dyn Optimizer<A, D>>,
553 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 pub fn new(inner: Box<dyn Optimizer<A, D>>, outer: Box<dyn Optimizer<A, D>>) -> Self {
569 Self { inner, outer }
570 }
571
572 pub fn inner(&self) -> &dyn Optimizer<A, D> {
574 self.inner.as_ref()
575 }
576
577 pub fn inner_mut(&mut self) -> &mut dyn Optimizer<A, D> {
579 self.inner.as_mut()
580 }
581
582 pub fn outer(&self) -> &dyn Optimizer<A, D> {
584 self.outer.as_ref()
585 }
586
587 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 let intermediate_params = self.inner.step(params, gradients)?;
601
602 self.outer.step(&intermediate_params, gradients)
604 }
605
606 fn get_learning_rate(&self) -> A {
607 self.inner.get_learning_rate()
609 }
610
611 fn set_learning_rate(&mut self, learningrate: A) {
612 self.inner.set_learning_rate(learningrate);
614 self.outer.set_learning_rate(learningrate);
615 }
616}
617
618pub struct WeightedOptimizer<A, D>
641where
642 A: Float + ScalarOperand + Debug,
643 D: Dimension,
644{
645 optimizers: Vec<Box<dyn Optimizer<A, D>>>,
647 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 pub fn new() -> Self {
668 Self {
669 optimizers: Vec::new(),
670 weights: Vec::new(),
671 }
672 }
673
674 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 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 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 pub fn num_optimizers(&self) -> usize {
711 self.optimizers.len()
712 }
713
714 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 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 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 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 let params = Array1::zeros(3);
797 let gradients = Array1::from_vec(vec![1.0, 2.0, 3.0]);
798
799 let updated_params = seq_optimizer
801 .step(¶ms, &gradients)
802 .expect("step succeeds in test_sequential_optimizer");
803
804 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 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); let group2 = ParameterGroup::new(params2.clone(), 1); 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 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 let updated_params = parallel_optimizer
833 .update_all_parameters(&[gradients1, gradients2])
834 .expect("update_all_parameters succeeds in test_parallel_optimizer");
835
836 assert_abs_diff_eq!(updated_params[0][0], -0.1);
839 assert_abs_diff_eq!(updated_params[0][1], -0.2);
840
841 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 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 let params = Array1::zeros(3);
859 let gradients = Array1::from_vec(vec![1.0, 2.0, 3.0]);
860
861 let updated_params = chained_optimizer
863 .step(¶ms, &gradients)
864 .expect("step succeeds in test_chained_optimizer");
865
866 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 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 assert_abs_diff_eq!(seq_optimizer.get_learning_rate(), 0.1);
885
886 seq_optimizer.set_learning_rate(0.05);
888
889 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 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 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 let params_refs = vec![¶ms1, ¶ms2, ¶ms3];
927 let gradients_refs = vec![&gradients1, &gradients2, &gradients3];
928
929 let updated_params = parallel_optimizer
930 .step_list(¶ms_refs, &gradients_refs)
931 .expect("step_list succeeds in test_parallel_optimizer_step_list");
932
933 assert_abs_diff_eq!(updated_params[0][0], -0.1);
936 assert_abs_diff_eq!(updated_params[0][1], -0.2);
937
938 assert!(updated_params[1][0] != -0.3);
941
942 assert!(updated_params[2][0] < 0.0);
945 }
946
947 #[test]
948 fn test_chained_optimizer_learning_rate() {
949 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 assert_abs_diff_eq!(chained_optimizer.get_learning_rate(), 0.1);
958
959 chained_optimizer.set_learning_rate(0.05);
961
962 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 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(¶ms, &gradients).expect("step failed");
983
984 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(¶ms, &gradients).expect("step failed");
1006
1007 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(¶ms, &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 assert_abs_diff_eq!(weighted.get_learning_rate(), 0.1);
1047
1048 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 #[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 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(&[¶ms1, ¶ms2], &[&grads1, &grads2])
1094 .expect("step_list failed");
1095
1096 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 #[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(&[¶ms], &[&grads]);
1129 assert!(result.is_err(), "empty optimizer list must be rejected");
1130 }
1131}