Skip to main content

torsh_optim/
composition.rs

1//! Optimizer composition tools
2//!
3//! This module provides utilities for composing multiple optimizers together,
4//! creating ensembles, pipelines, and adaptive switching between optimizers.
5
6use 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/// Strategies for combining multiple optimizers
15#[derive(Debug, Clone)]
16pub enum CompositionStrategy {
17    /// Use different optimizers in sequence (pipeline)
18    Sequential {
19        schedule: Vec<(String, usize)>, // (optimizer_name, steps)
20    },
21    /// Ensemble of optimizers with weighted averaging
22    Ensemble {
23        weights: HashMap<String, f32>,
24        combination_method: CombinationMethod,
25    },
26    /// Adaptive switching based on performance
27    Adaptive {
28        switch_criterion: SwitchCriterion,
29        evaluation_window: usize,
30    },
31    /// Parallel execution with consensus
32    Consensus {
33        agreement_threshold: f32,
34        voting_method: VotingMethod,
35    },
36    /// Hierarchical composition
37    Hierarchical { levels: Vec<CompositionLevel> },
38}
39
40#[derive(Debug, Clone)]
41pub enum CombinationMethod {
42    /// Weighted average of parameter updates
43    WeightedAverage,
44    /// Median of parameter updates
45    Median,
46    /// Best performing optimizer takes precedence
47    BestWins,
48    /// Custom combination function
49    Custom(fn(&[Tensor]) -> Tensor),
50}
51
52#[derive(Debug, Clone)]
53pub enum SwitchCriterion {
54    /// Switch based on loss improvement
55    LossImprovement { threshold: f32 },
56    /// Switch based on gradient magnitude
57    GradientMagnitude { threshold: f32 },
58    /// Switch based on convergence rate
59    ConvergenceRate { window: usize },
60    /// Switch based on custom metric
61    Custom(fn(&OptimizerMetrics) -> bool),
62}
63
64#[derive(Debug, Clone)]
65pub enum VotingMethod {
66    /// Majority vote on update direction
67    Majority,
68    /// Weighted vote based on recent performance
69    WeightedVote,
70    /// Unanimous consensus required
71    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/// Metrics for optimizer performance evaluation
82#[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        // Compute convergence rate
124        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        // Compute stability score (lower variance is better)
132        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        // Compute efficiency score (convergence per unit work)
145        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
154/// Composed optimizer that combines multiple optimizers
155pub 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    /// Add an optimizer to the composition
216    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    /// Remove an optimizer from the composition
222    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    /// Get current active optimizers
228    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                // Return optimizers from current level
252                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    /// Update optimizer metrics
269    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    /// Get performance metrics for an optimizer
282    pub fn get_metrics(&self, optimizer_name: &str) -> Option<&OptimizerMetrics> {
283        self.metrics.get(optimizer_name)
284    }
285
286    /// Get the best performing optimizer based on metrics
287    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    /// Execute the composition strategy
304    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                    // Continue with current optimizer
334                    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                    // Move to next phase
341                    *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        // Store parameter updates from each optimizer
360        let mut updates = HashMap::new();
361
362        // Collect updates first to avoid borrowing conflicts
363        let mut state_diffs = Vec::new();
364        for (name, optimizer) in &mut self.optimizers {
365            if weights.contains_key(name) {
366                // Get parameters before step
367                let state_before = optimizer.state_dict()?;
368
369                // Perform optimizer step
370                optimizer.step()?;
371
372                // Get parameters after step
373                let state_after = optimizer.state_dict()?;
374
375                state_diffs.push((name.clone(), state_before, state_after));
376            }
377        }
378
379        // Compute updates after collecting all state diffs
380        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        // Combine updates according to strategy
386        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        // Evaluate current optimizer performance
400        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                // Check if we should switch
405                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 = &current_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            // No current optimizer, select the best one
440            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        // Collect votes (parameter updates) from all optimizers
452        let mut votes = HashMap::new();
453
454        // Collect state diffs first to avoid borrowing conflicts
455        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        // Compute updates after collecting all state diffs
465        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        // Apply voting mechanism
471        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        // Get current level without borrowing the entire state
486        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            // Execute the strategy for the current level
497            // This would recursively handle the composition strategy for this level
498            // For now, just execute the first optimizer in the level
499            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            // Check if we should move to the next level
506            // This would be based on some criteria (e.g., convergence, time, etc.)
507            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    // Helper methods
520    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        // Compute difference for each parameter
528        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        // Get all parameter names
552        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        // Combine each parameter
560        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(&param_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(&param_updates, &param_weights)?
575                    }
576                    CombinationMethod::Median => self.median_update(&param_updates)?,
577                    CombinationMethod::BestWins => {
578                        // Use update from best performing optimizer
579                        param_updates[0].clone() // Simplified
580                    }
581                    CombinationMethod::Custom(combine_fn) => combine_fn(&param_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        // For simplicity, just return the middle tensor
626        // In practice, would compute element-wise median
627        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        // Simplified majority vote - just average all votes
636        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        // Average by count
652        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        // Weight votes by optimizer performance
668        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        // Apply weighted combination
679        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        // Only apply updates where optimizers agree (within threshold)
688        let mut unanimous_updates = HashMap::new();
689
690        // Get all parameter names
691        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(&param_name) {
703                    param_votes.push(param_tensor.clone());
704                }
705            }
706
707            if param_votes.len() > 1 {
708                // Check agreement (simplified - use variance)
709                let mean = self.compute_mean_tensor(&param_votes)?;
710                let variance = self.compute_variance_tensor(&param_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        // Apply the combined update to the actual parameters
761        // This would require access to the actual model parameters
762        // For now, this is a placeholder
763        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        // Placeholder logic for hierarchical advancement
777        // In practice, this would check convergence criteria, time limits, etc.
778        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        // Return learning rates from all optimizers
796        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 there is an active/current optimizer, return its parameters.
817        if let Some(current) = &self.current_optimizer {
818            if let Some(optimizer) = self.optimizers.get(current) {
819                return optimizer.parameters();
820            }
821        }
822
823        // Otherwise, collect parameters from all composed optimizers.
824        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        // Combine state from all optimizers
833        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            // Add optimizer name prefix to avoid conflicts
840            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        // Split state back to individual optimizers
859        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
888/// Builder for composed optimizers
889pub 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
932/// Utility functions for optimizer composition
933pub mod utils {
934    use super::*;
935
936    /// Create an ensemble of optimizers with equal weights
937    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    /// Create a sequential pipeline of optimizers
958    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    /// Create an adaptive composition that switches based on loss improvement
977    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}