1use crate::{Optimizer, OptimizerError, OptimizerResult, OptimizerState};
7use parking_lot::RwLock;
8use serde::{Deserialize, Serialize};
9use std::collections::HashMap;
10use std::ops::Add;
11use std::sync::Arc;
12use torsh_tensor::Tensor;
13
14#[derive(Debug, Clone)]
16pub enum CompositionStrategy {
17 Sequential {
19 schedule: Vec<(String, usize)>, },
21 Ensemble {
23 weights: HashMap<String, f32>,
24 combination_method: CombinationMethod,
25 },
26 Adaptive {
28 switch_criterion: SwitchCriterion,
29 evaluation_window: usize,
30 },
31 Consensus {
33 agreement_threshold: f32,
34 voting_method: VotingMethod,
35 },
36 Hierarchical { levels: Vec<CompositionLevel> },
38}
39
40#[derive(Debug, Clone)]
41pub enum CombinationMethod {
42 WeightedAverage,
44 Median,
46 BestWins,
48 Custom(fn(&[Tensor]) -> Tensor),
50}
51
52#[derive(Debug, Clone)]
53pub enum SwitchCriterion {
54 LossImprovement { threshold: f32 },
56 GradientMagnitude { threshold: f32 },
58 ConvergenceRate { window: usize },
60 Custom(fn(&OptimizerMetrics) -> bool),
62}
63
64#[derive(Debug, Clone)]
65pub enum VotingMethod {
66 Majority,
68 WeightedVote,
70 Unanimous,
72}
73
74#[derive(Debug, Clone)]
75pub struct CompositionLevel {
76 pub name: String,
77 pub optimizers: Vec<String>,
78 pub strategy: CompositionStrategy,
79}
80
81#[derive(Debug, Clone, Serialize, Deserialize)]
83pub struct OptimizerMetrics {
84 pub loss_history: Vec<f32>,
85 pub gradient_norms: Vec<f32>,
86 pub update_magnitudes: Vec<f32>,
87 pub convergence_rate: f32,
88 pub stability_score: f32,
89 pub efficiency_score: f32,
90}
91
92impl Default for OptimizerMetrics {
93 fn default() -> Self {
94 Self {
95 loss_history: Vec::new(),
96 gradient_norms: Vec::new(),
97 update_magnitudes: Vec::new(),
98 convergence_rate: 0.0,
99 stability_score: 0.0,
100 efficiency_score: 0.0,
101 }
102 }
103}
104
105impl OptimizerMetrics {
106 pub fn new() -> Self {
107 Self::default()
108 }
109
110 pub fn update(&mut self, loss: f32, gradient_norm: f32, update_magnitude: f32) {
111 self.loss_history.push(loss);
112 self.gradient_norms.push(gradient_norm);
113 self.update_magnitudes.push(update_magnitude);
114
115 self.compute_derived_metrics();
116 }
117
118 fn compute_derived_metrics(&mut self) {
119 if self.loss_history.len() < 2 {
120 return;
121 }
122
123 let recent_losses = &self.loss_history[self.loss_history.len().saturating_sub(10)..];
125 if recent_losses.len() >= 2 {
126 let start_loss = recent_losses[0];
127 let end_loss = recent_losses[recent_losses.len() - 1];
128 self.convergence_rate = (start_loss - end_loss) / recent_losses.len() as f32;
129 }
130
131 if !self.update_magnitudes.is_empty() {
133 let mean_magnitude =
134 self.update_magnitudes.iter().sum::<f32>() / self.update_magnitudes.len() as f32;
135 let variance = self
136 .update_magnitudes
137 .iter()
138 .map(|x| (x - mean_magnitude).powi(2))
139 .sum::<f32>()
140 / self.update_magnitudes.len() as f32;
141 self.stability_score = 1.0 / (1.0 + variance);
142 }
143
144 if !self.loss_history.is_empty() {
146 let total_improvement =
147 self.loss_history[0] - self.loss_history[self.loss_history.len() - 1];
148 let steps = self.loss_history.len() as f32;
149 self.efficiency_score = total_improvement / steps;
150 }
151 }
152}
153
154pub struct ComposedOptimizer {
156 strategy: CompositionStrategy,
157 optimizers: HashMap<String, Box<dyn Optimizer>>,
158 metrics: HashMap<String, OptimizerMetrics>,
159 current_optimizer: Option<String>,
160 step_count: usize,
161 composition_state: CompositionState,
162}
163
164#[derive(Debug, Clone)]
165enum CompositionState {
166 Sequential {
167 current_phase: usize,
168 phase_steps: usize,
169 },
170 Ensemble {
171 last_updates: HashMap<String, Vec<Tensor>>,
172 },
173 Adaptive {
174 evaluation_buffer: Vec<(String, f32)>,
175 },
176 Consensus {
177 votes: HashMap<String, Vec<Tensor>>,
178 },
179 Hierarchical {
180 current_level: usize,
181 },
182}
183
184impl ComposedOptimizer {
185 pub fn new(strategy: CompositionStrategy) -> Self {
186 let composition_state = match &strategy {
187 CompositionStrategy::Sequential { .. } => CompositionState::Sequential {
188 current_phase: 0,
189 phase_steps: 0,
190 },
191 CompositionStrategy::Ensemble { .. } => CompositionState::Ensemble {
192 last_updates: HashMap::new(),
193 },
194 CompositionStrategy::Adaptive { .. } => CompositionState::Adaptive {
195 evaluation_buffer: Vec::new(),
196 },
197 CompositionStrategy::Consensus { .. } => CompositionState::Consensus {
198 votes: HashMap::new(),
199 },
200 CompositionStrategy::Hierarchical { .. } => {
201 CompositionState::Hierarchical { current_level: 0 }
202 }
203 };
204
205 Self {
206 strategy,
207 optimizers: HashMap::new(),
208 metrics: HashMap::new(),
209 current_optimizer: None,
210 step_count: 0,
211 composition_state,
212 }
213 }
214
215 pub fn add_optimizer(&mut self, name: String, optimizer: Box<dyn Optimizer>) {
217 self.optimizers.insert(name.clone(), optimizer);
218 self.metrics.insert(name.clone(), OptimizerMetrics::new());
219 }
220
221 pub fn remove_optimizer(&mut self, name: &str) -> Option<Box<dyn Optimizer>> {
223 self.metrics.remove(name);
224 self.optimizers.remove(name)
225 }
226
227 pub fn active_optimizers(&self) -> Vec<&str> {
229 match &self.strategy {
230 CompositionStrategy::Sequential { .. } => {
231 if let Some(ref current) = self.current_optimizer {
232 vec![current]
233 } else {
234 Vec::new()
235 }
236 }
237 CompositionStrategy::Ensemble { .. } => {
238 self.optimizers.keys().map(|s| s.as_str()).collect()
239 }
240 CompositionStrategy::Adaptive { .. } => {
241 if let Some(ref current) = self.current_optimizer {
242 vec![current]
243 } else {
244 Vec::new()
245 }
246 }
247 CompositionStrategy::Consensus { .. } => {
248 self.optimizers.keys().map(|s| s.as_str()).collect()
249 }
250 CompositionStrategy::Hierarchical { .. } => {
251 if let CompositionState::Hierarchical { current_level } = &self.composition_state {
253 if let CompositionStrategy::Hierarchical { levels } = &self.strategy {
254 if *current_level < levels.len() {
255 return levels[*current_level]
256 .optimizers
257 .iter()
258 .map(|s| s.as_str())
259 .collect();
260 }
261 }
262 }
263 Vec::new()
264 }
265 }
266 }
267
268 pub fn update_metrics(
270 &mut self,
271 optimizer_name: &str,
272 loss: f32,
273 gradient_norm: f32,
274 update_magnitude: f32,
275 ) {
276 if let Some(metrics) = self.metrics.get_mut(optimizer_name) {
277 metrics.update(loss, gradient_norm, update_magnitude);
278 }
279 }
280
281 pub fn get_metrics(&self, optimizer_name: &str) -> Option<&OptimizerMetrics> {
283 self.metrics.get(optimizer_name)
284 }
285
286 pub fn best_optimizer(&self) -> Option<&str> {
288 let mut best_name = None;
289 let mut best_score = f32::NEG_INFINITY;
290
291 for (name, metrics) in &self.metrics {
292 let score =
293 metrics.convergence_rate * metrics.stability_score * metrics.efficiency_score;
294 if score > best_score {
295 best_score = score;
296 best_name = Some(name.as_str());
297 }
298 }
299
300 best_name
301 }
302
303 fn execute_strategy(&mut self) -> OptimizerResult<()> {
305 match &self.strategy.clone() {
306 CompositionStrategy::Sequential { schedule } => self.execute_sequential(schedule),
307 CompositionStrategy::Ensemble {
308 weights,
309 combination_method,
310 } => self.execute_ensemble(weights, combination_method),
311 CompositionStrategy::Adaptive {
312 switch_criterion,
313 evaluation_window,
314 } => self.execute_adaptive(switch_criterion, *evaluation_window),
315 CompositionStrategy::Consensus {
316 agreement_threshold,
317 voting_method,
318 } => self.execute_consensus(*agreement_threshold, voting_method),
319 CompositionStrategy::Hierarchical { levels } => self.execute_hierarchical(levels),
320 }
321 }
322
323 fn execute_sequential(&mut self, schedule: &[(String, usize)]) -> OptimizerResult<()> {
324 if let CompositionState::Sequential {
325 current_phase,
326 phase_steps,
327 } = &mut self.composition_state
328 {
329 if *current_phase < schedule.len() {
330 let (optimizer_name, max_steps) = &schedule[*current_phase];
331
332 if *phase_steps < *max_steps {
333 if let Some(optimizer) = self.optimizers.get_mut(optimizer_name) {
335 optimizer.step()?;
336 *phase_steps += 1;
337 self.current_optimizer = Some(optimizer_name.clone());
338 }
339 } else {
340 *current_phase += 1;
342 *phase_steps = 0;
343
344 if *current_phase < schedule.len() {
345 let (next_optimizer, _) = &schedule[*current_phase];
346 self.current_optimizer = Some(next_optimizer.clone());
347 }
348 }
349 }
350 }
351 Ok(())
352 }
353
354 fn execute_ensemble(
355 &mut self,
356 weights: &HashMap<String, f32>,
357 combination_method: &CombinationMethod,
358 ) -> OptimizerResult<()> {
359 let mut updates = HashMap::new();
361
362 let mut state_diffs = Vec::new();
364 for (name, optimizer) in &mut self.optimizers {
365 if weights.contains_key(name) {
366 let state_before = optimizer.state_dict()?;
368
369 optimizer.step()?;
371
372 let state_after = optimizer.state_dict()?;
374
375 state_diffs.push((name.clone(), state_before, state_after));
376 }
377 }
378
379 for (name, state_before, state_after) in state_diffs {
381 let update = self.compute_parameter_update(&state_before, &state_after)?;
382 updates.insert(name, update);
383 }
384
385 if !updates.is_empty() {
387 let combined_update = self.combine_updates(&updates, weights, combination_method)?;
388 self.apply_combined_update(&combined_update)?;
389 }
390
391 Ok(())
392 }
393
394 fn execute_adaptive(
395 &mut self,
396 switch_criterion: &SwitchCriterion,
397 evaluation_window: usize,
398 ) -> OptimizerResult<()> {
399 if let Some(current_name) = &self.current_optimizer.clone() {
401 if let Some(current_optimizer) = self.optimizers.get_mut(current_name) {
402 current_optimizer.step()?;
403
404 if let Some(current_metrics) = self.metrics.get(current_name) {
406 let should_switch = match switch_criterion {
407 SwitchCriterion::LossImprovement { threshold } => {
408 if current_metrics.loss_history.len() >= evaluation_window {
409 let recent_losses = ¤t_metrics.loss_history
410 [current_metrics.loss_history.len() - evaluation_window..];
411 let improvement =
412 recent_losses[0] - recent_losses[recent_losses.len() - 1];
413 improvement < *threshold
414 } else {
415 false
416 }
417 }
418 SwitchCriterion::GradientMagnitude { threshold } => {
419 if let Some(&last_gradient_norm) = current_metrics.gradient_norms.last()
420 {
421 last_gradient_norm < *threshold
422 } else {
423 false
424 }
425 }
426 SwitchCriterion::ConvergenceRate { window } => {
427 current_metrics.convergence_rate < 0.001
428 && current_metrics.loss_history.len() >= *window
429 }
430 SwitchCriterion::Custom(criterion_fn) => criterion_fn(current_metrics),
431 };
432
433 if should_switch {
434 self.switch_to_best_optimizer()?;
435 }
436 }
437 }
438 } else {
439 self.switch_to_best_optimizer()?;
441 }
442
443 Ok(())
444 }
445
446 fn execute_consensus(
447 &mut self,
448 agreement_threshold: f32,
449 voting_method: &VotingMethod,
450 ) -> OptimizerResult<()> {
451 let mut votes = HashMap::new();
453
454 let mut state_diffs = Vec::new();
456 for (name, optimizer) in &mut self.optimizers {
457 let state_before = optimizer.state_dict()?;
458 optimizer.step()?;
459 let state_after = optimizer.state_dict()?;
460
461 state_diffs.push((name.clone(), state_before, state_after));
462 }
463
464 for (name, state_before, state_after) in state_diffs {
466 let update = self.compute_parameter_update(&state_before, &state_after)?;
467 votes.insert(name, update);
468 }
469
470 if !votes.is_empty() {
472 let consensus_update = match voting_method {
473 VotingMethod::Majority => self.majority_vote(&votes)?,
474 VotingMethod::WeightedVote => self.weighted_vote(&votes)?,
475 VotingMethod::Unanimous => self.unanimous_vote(&votes, agreement_threshold)?,
476 };
477
478 self.apply_combined_update(&consensus_update)?;
479 }
480
481 Ok(())
482 }
483
484 fn execute_hierarchical(&mut self, levels: &[CompositionLevel]) -> OptimizerResult<()> {
485 let current_level_idx =
487 if let CompositionState::Hierarchical { current_level } = &self.composition_state {
488 *current_level
489 } else {
490 return Ok(());
491 };
492
493 if current_level_idx < levels.len() {
494 let level = &levels[current_level_idx];
495
496 if let Some(optimizer_name) = level.optimizers.first() {
500 if let Some(optimizer) = self.optimizers.get_mut(optimizer_name) {
501 optimizer.step()?;
502 }
503 }
504
505 if self.should_advance_level() {
508 if let CompositionState::Hierarchical { current_level } =
509 &mut self.composition_state
510 {
511 *current_level += 1;
512 }
513 }
514 }
515
516 Ok(())
517 }
518
519 fn compute_parameter_update(
521 &self,
522 state_before: &OptimizerState,
523 state_after: &OptimizerState,
524 ) -> OptimizerResult<HashMap<String, Tensor>> {
525 let mut updates = HashMap::new();
526
527 for (param_name, param_dict_after) in &state_after.state {
529 if let Some(param_dict_before) = state_before.state.get(param_name) {
530 for (state_name, tensor_after) in param_dict_after {
531 if let Some(tensor_before) = param_dict_before.get(state_name) {
532 let update = tensor_after.sub(tensor_before)?;
533 let full_name = format!("{param_name}_{state_name}");
534 updates.insert(full_name, update);
535 }
536 }
537 }
538 }
539
540 Ok(updates)
541 }
542
543 fn combine_updates(
544 &self,
545 updates: &HashMap<String, HashMap<String, Tensor>>,
546 weights: &HashMap<String, f32>,
547 combination_method: &CombinationMethod,
548 ) -> OptimizerResult<HashMap<String, Tensor>> {
549 let mut combined = HashMap::new();
550
551 let mut all_param_names = std::collections::HashSet::new();
553 for update_dict in updates.values() {
554 for param_name in update_dict.keys() {
555 all_param_names.insert(param_name.clone());
556 }
557 }
558
559 for param_name in all_param_names {
561 let mut param_updates = Vec::new();
562 let mut param_weights = Vec::new();
563
564 for (optimizer_name, update_dict) in updates {
565 if let Some(param_update) = update_dict.get(¶m_name) {
566 param_updates.push(param_update.clone());
567 param_weights.push(weights.get(optimizer_name).copied().unwrap_or(1.0));
568 }
569 }
570
571 if !param_updates.is_empty() {
572 let combined_update = match combination_method {
573 CombinationMethod::WeightedAverage => {
574 self.weighted_average(¶m_updates, ¶m_weights)?
575 }
576 CombinationMethod::Median => self.median_update(¶m_updates)?,
577 CombinationMethod::BestWins => {
578 param_updates[0].clone() }
581 CombinationMethod::Custom(combine_fn) => combine_fn(¶m_updates),
582 };
583
584 combined.insert(param_name, combined_update);
585 }
586 }
587
588 Ok(combined)
589 }
590
591 fn weighted_average(&self, tensors: &[Tensor], weights: &[f32]) -> OptimizerResult<Tensor> {
592 if tensors.is_empty() || weights.is_empty() || tensors.len() != weights.len() {
593 return Err(OptimizerError::InvalidParameter(
594 "Mismatched tensors and weights".to_string(),
595 ));
596 }
597
598 let weight_sum: f32 = weights.iter().sum();
599 if weight_sum == 0.0 {
600 return Err(OptimizerError::InvalidParameter(
601 "Zero weight sum".to_string(),
602 ));
603 }
604
605 let mut result = tensors[0].mul_scalar(weights[0] / weight_sum)?;
606 for i in 1..tensors.len() {
607 let weighted_tensor = tensors[i].mul_scalar(weights[i] / weight_sum)?;
608 result = result.add(&weighted_tensor)?;
609 }
610
611 Ok(result)
612 }
613
614 fn median_update(&self, tensors: &[Tensor]) -> OptimizerResult<Tensor> {
615 if tensors.is_empty() {
616 return Err(OptimizerError::InvalidParameter(
617 "Empty tensor list".to_string(),
618 ));
619 }
620
621 if tensors.len() == 1 {
622 return Ok(tensors[0].clone());
623 }
624
625 let median_idx = tensors.len() / 2;
628 Ok(tensors[median_idx].clone())
629 }
630
631 fn majority_vote(
632 &self,
633 votes: &HashMap<String, HashMap<String, Tensor>>,
634 ) -> OptimizerResult<HashMap<String, Tensor>> {
635 let mut combined = HashMap::new();
637 let mut param_counts = HashMap::new();
638
639 for vote_dict in votes.values() {
640 for (param_name, param_tensor) in vote_dict {
641 combined
642 .entry(param_name.clone())
643 .and_modify(|t: &mut Tensor| {
644 *t = t.add(param_tensor).expect("tensor add should succeed")
645 })
646 .or_insert(param_tensor.clone());
647 *param_counts.entry(param_name.clone()).or_insert(0) += 1;
648 }
649 }
650
651 for (param_name, tensor) in &mut combined {
653 if let Some(&count) = param_counts.get(param_name) {
654 if count > 1 {
655 *tensor = tensor.div_scalar(count as f32)?;
656 }
657 }
658 }
659
660 Ok(combined)
661 }
662
663 fn weighted_vote(
664 &self,
665 votes: &HashMap<String, HashMap<String, Tensor>>,
666 ) -> OptimizerResult<HashMap<String, Tensor>> {
667 let mut weights = HashMap::new();
669 for optimizer_name in votes.keys() {
670 if let Some(metrics) = self.metrics.get(optimizer_name) {
671 let weight = metrics.efficiency_score * metrics.stability_score;
672 weights.insert(optimizer_name.clone(), weight);
673 } else {
674 weights.insert(optimizer_name.clone(), 1.0);
675 }
676 }
677
678 self.combine_updates(votes, &weights, &CombinationMethod::WeightedAverage)
680 }
681
682 fn unanimous_vote(
683 &self,
684 votes: &HashMap<String, HashMap<String, Tensor>>,
685 agreement_threshold: f32,
686 ) -> OptimizerResult<HashMap<String, Tensor>> {
687 let mut unanimous_updates = HashMap::new();
689
690 let mut all_params = std::collections::HashSet::new();
692 for vote_dict in votes.values() {
693 for param_name in vote_dict.keys() {
694 all_params.insert(param_name.clone());
695 }
696 }
697
698 for param_name in all_params {
699 let mut param_votes = Vec::new();
700
701 for vote_dict in votes.values() {
702 if let Some(param_tensor) = vote_dict.get(¶m_name) {
703 param_votes.push(param_tensor.clone());
704 }
705 }
706
707 if param_votes.len() > 1 {
708 let mean = self.compute_mean_tensor(¶m_votes)?;
710 let variance = self.compute_variance_tensor(¶m_votes, &mean)?;
711 let variance_norm = variance.norm()?.item()?;
712
713 if variance_norm < agreement_threshold {
714 unanimous_updates.insert(param_name, mean);
715 }
716 } else if param_votes.len() == 1 {
717 unanimous_updates.insert(param_name, param_votes[0].clone());
718 }
719 }
720
721 Ok(unanimous_updates)
722 }
723
724 fn compute_mean_tensor(&self, tensors: &[Tensor]) -> OptimizerResult<Tensor> {
725 if tensors.is_empty() {
726 return Err(OptimizerError::InvalidParameter(
727 "Empty tensor list".to_string(),
728 ));
729 }
730
731 let mut sum = tensors[0].clone();
732 for tensor in tensors.iter().skip(1) {
733 sum = sum.add(tensor)?;
734 }
735
736 Ok(sum.div_scalar(tensors.len() as f32)?)
737 }
738
739 fn compute_variance_tensor(
740 &self,
741 tensors: &[Tensor],
742 mean: &Tensor,
743 ) -> OptimizerResult<Tensor> {
744 if tensors.is_empty() {
745 return Err(OptimizerError::InvalidParameter(
746 "Empty tensor list".to_string(),
747 ));
748 }
749
750 let mut variance = tensors[0].sub(mean)?.pow_scalar(2.0)?;
751 for tensor in tensors.iter().skip(1) {
752 let diff = tensor.sub(mean)?.pow_scalar(2.0)?;
753 variance = variance.add(&diff)?;
754 }
755
756 Ok(variance.div_scalar(tensors.len() as f32)?)
757 }
758
759 fn apply_combined_update(&mut self, _update: &HashMap<String, Tensor>) -> OptimizerResult<()> {
760 Ok(())
764 }
765
766 fn switch_to_best_optimizer(&mut self) -> OptimizerResult<()> {
767 let best_name = self.best_optimizer().map(|s| s.to_string());
768 if let Some(best_name) = best_name {
769 self.current_optimizer = Some(best_name.clone());
770 log::info!("Switched to optimizer: {best_name}");
771 }
772 Ok(())
773 }
774
775 fn should_advance_level(&self) -> bool {
776 false
779 }
780}
781
782impl Optimizer for ComposedOptimizer {
783 fn step(&mut self) -> OptimizerResult<()> {
784 self.step_count += 1;
785 self.execute_strategy()
786 }
787
788 fn zero_grad(&mut self) {
789 for optimizer in self.optimizers.values_mut() {
790 optimizer.zero_grad();
791 }
792 }
793
794 fn get_lr(&self) -> Vec<f32> {
795 let mut all_lrs = Vec::new();
797 for optimizer in self.optimizers.values() {
798 all_lrs.extend(optimizer.get_lr());
799 }
800 all_lrs
801 }
802
803 fn set_lr(&mut self, lr: f32) {
804 for optimizer in self.optimizers.values_mut() {
805 optimizer.set_lr(lr);
806 }
807 }
808
809 fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
810 for optimizer in self.optimizers.values_mut() {
811 optimizer.add_param_group(params.clone(), options.clone());
812 }
813 }
814
815 fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
816 if let Some(current) = &self.current_optimizer {
818 if let Some(optimizer) = self.optimizers.get(current) {
819 return optimizer.parameters();
820 }
821 }
822
823 let mut all_params = Vec::new();
825 for optimizer in self.optimizers.values() {
826 all_params.extend(optimizer.parameters());
827 }
828 all_params
829 }
830
831 fn state_dict(&self) -> OptimizerResult<OptimizerState> {
832 let mut combined_state = HashMap::new();
834 let mut combined_param_groups = Vec::new();
835
836 for (name, optimizer) in &self.optimizers {
837 let state = optimizer.state_dict()?;
838
839 for (param_id, param_state) in state.state {
841 let prefixed_id = format!("{name}_{param_id}");
842 combined_state.insert(prefixed_id, param_state);
843 }
844
845 combined_param_groups.extend(state.param_groups);
846 }
847
848 Ok(OptimizerState {
849 optimizer_type: "CompositeOptimizer".to_string(),
850 version: "0.1.0".to_string(),
851 param_groups: combined_param_groups,
852 state: combined_state,
853 global_state: HashMap::new(),
854 })
855 }
856
857 fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
858 for (optimizer_name, optimizer) in &mut self.optimizers {
860 let mut optimizer_state = HashMap::new();
861 let prefix = format!("{optimizer_name}_");
862
863 for (param_id, param_state) in &state.state {
864 if param_id.starts_with(&prefix) {
865 let unprefixed_id = param_id
866 .strip_prefix(&prefix)
867 .expect("prefix should exist after starts_with check")
868 .to_string();
869 optimizer_state.insert(unprefixed_id, param_state.clone());
870 }
871 }
872
873 let optimizer_state_dict = OptimizerState {
874 optimizer_type: "CompositeOptimizer".to_string(),
875 version: "0.1.0".to_string(),
876 param_groups: state.param_groups.clone(),
877 state: optimizer_state,
878 global_state: HashMap::new(),
879 };
880
881 optimizer.load_state_dict(optimizer_state_dict)?;
882 }
883
884 Ok(())
885 }
886}
887
888pub struct CompositionBuilder {
890 strategy: Option<CompositionStrategy>,
891 optimizers: HashMap<String, Box<dyn Optimizer>>,
892}
893
894impl Default for CompositionBuilder {
895 fn default() -> Self {
896 Self {
897 strategy: None,
898 optimizers: HashMap::new(),
899 }
900 }
901}
902
903impl CompositionBuilder {
904 pub fn new() -> Self {
905 Self::default()
906 }
907
908 pub fn strategy(mut self, strategy: CompositionStrategy) -> Self {
909 self.strategy = Some(strategy);
910 self
911 }
912
913 pub fn add_optimizer(mut self, name: &str, optimizer: Box<dyn Optimizer>) -> Self {
914 self.optimizers.insert(name.to_string(), optimizer);
915 self
916 }
917
918 pub fn build(self) -> OptimizerResult<ComposedOptimizer> {
919 let strategy = self.strategy.ok_or_else(|| {
920 OptimizerError::ConfigError("No composition strategy specified".to_string())
921 })?;
922
923 let mut composed = ComposedOptimizer::new(strategy);
924 for (name, optimizer) in self.optimizers {
925 composed.add_optimizer(name, optimizer);
926 }
927
928 Ok(composed)
929 }
930}
931
932pub mod utils {
934 use super::*;
935
936 pub fn equal_ensemble(
938 optimizers: Vec<(&str, Box<dyn Optimizer>)>,
939 ) -> OptimizerResult<ComposedOptimizer> {
940 let mut weights = HashMap::new();
941 let weight = 1.0 / optimizers.len() as f32;
942
943 let mut builder = CompositionBuilder::new();
944 for (name, optimizer) in optimizers {
945 weights.insert(name.to_string(), weight);
946 builder = builder.add_optimizer(name, optimizer);
947 }
948
949 let strategy = CompositionStrategy::Ensemble {
950 weights,
951 combination_method: CombinationMethod::WeightedAverage,
952 };
953
954 builder.strategy(strategy).build()
955 }
956
957 pub fn sequential_pipeline(
959 schedule: Vec<(&str, Box<dyn Optimizer>, usize)>,
960 ) -> OptimizerResult<ComposedOptimizer> {
961 let mut builder = CompositionBuilder::new();
962 let mut strategy_schedule = Vec::new();
963
964 for (name, optimizer, steps) in schedule {
965 builder = builder.add_optimizer(name, optimizer);
966 strategy_schedule.push((name.to_string(), steps));
967 }
968
969 let strategy = CompositionStrategy::Sequential {
970 schedule: strategy_schedule,
971 };
972
973 builder.strategy(strategy).build()
974 }
975
976 pub fn adaptive_switching(
978 optimizers: Vec<(&str, Box<dyn Optimizer>)>,
979 improvement_threshold: f32,
980 ) -> OptimizerResult<ComposedOptimizer> {
981 let mut builder = CompositionBuilder::new();
982
983 for (name, optimizer) in optimizers {
984 builder = builder.add_optimizer(name, optimizer);
985 }
986
987 let strategy = CompositionStrategy::Adaptive {
988 switch_criterion: SwitchCriterion::LossImprovement {
989 threshold: improvement_threshold,
990 },
991 evaluation_window: 10,
992 };
993
994 builder.strategy(strategy).build()
995 }
996}