Skip to main content

torsh_optim/
green_ai.rs

1//! Green AI Optimizers
2//!
3//! This module implements energy-aware and environmentally-conscious optimization
4//! algorithms that minimize computational carbon footprint while maintaining
5//! model performance.
6//!
7//! # Key Concepts
8//!
9//! - **Energy Efficiency**: Minimizing computational energy consumption
10//! - **Carbon Footprint**: Tracking and reducing CO2 emissions from training
11//! - **Adaptive Precision**: Dynamic adjustment of computation precision
12//! - **Sparse Computation**: Selective parameter updates to reduce FLOPs
13//!
14//! # Algorithms
15//!
16//! ## Energy-Aware Optimizer
17//!
18//! Tracks energy consumption per step and adapts learning rate based on
19//! energy efficiency metrics. Implements early stopping based on energy budgets.
20//!
21//! ## Carbon-Conscious Optimizer
22//!
23//! Monitors carbon emissions during training using real-time grid carbon intensity
24//! data. Schedules computationally intensive operations during low-carbon periods.
25//!
26//! ## Adaptive Precision Optimizer
27//!
28//! Dynamically adjusts numerical precision (FP32 ↔ FP16 ↔ BF16) based on
29//! gradient magnitudes and training stability to reduce energy consumption.
30//!
31//! ## Power-Capped Optimizer
32//!
33//! Enforces power consumption limits by adjusting batch sizes and learning rates
34//! to stay within specified power budgets.
35//!
36//! # References
37//!
38//! - Schwartz et al. (2020). "Green AI"
39//! - Strubell et al. (2019). "Energy and Policy Considerations for Deep Learning in NLP"
40//! - Patterson et al. (2021). "Carbon Emissions and Large Neural Network Training"
41//! - Lacoste et al. (2019). "Quantifying the Carbon Emissions of Machine Learning"
42
43use crate::{Optimizer, OptimizerError, OptimizerResult, OptimizerState};
44use parking_lot::RwLock;
45use std::collections::HashMap;
46use std::sync::Arc;
47use std::time::{Duration, Instant};
48use torsh_tensor::Tensor;
49
50// ============================================================================
51// Energy Tracking
52// ============================================================================
53
54/// Energy consumption metrics
55#[derive(Debug, Clone)]
56pub struct EnergyMetrics {
57    /// Total energy consumed (kWh)
58    pub total_energy_kwh: f64,
59    /// Average power (Watts)
60    pub avg_power_watts: f64,
61    /// Peak power (Watts)
62    pub peak_power_watts: f64,
63    /// Number of steps
64    pub num_steps: usize,
65    /// Total training time
66    pub total_time: Duration,
67    /// Energy per step (Joules)
68    pub energy_per_step: f64,
69}
70
71impl Default for EnergyMetrics {
72    fn default() -> Self {
73        Self {
74            total_energy_kwh: 0.0,
75            avg_power_watts: 0.0,
76            peak_power_watts: 0.0,
77            num_steps: 0,
78            total_time: Duration::from_secs(0),
79            energy_per_step: 0.0,
80        }
81    }
82}
83
84/// Carbon intensity (gCO2/kWh) - varies by location and time
85#[derive(Debug, Clone)]
86pub struct CarbonIntensity {
87    /// Current grid carbon intensity (gCO2/kWh)
88    pub intensity: f64,
89    /// Location/region
90    pub region: String,
91}
92
93impl Default for CarbonIntensity {
94    fn default() -> Self {
95        Self {
96            intensity: 475.0, // Global average
97            region: "global".to_string(),
98        }
99    }
100}
101
102// ============================================================================
103// Energy-Aware Optimizer
104// ============================================================================
105
106/// Energy-aware configuration
107#[derive(Debug, Clone)]
108pub struct EnergyAwareConfig {
109    /// Energy budget (kWh)
110    pub energy_budget_kwh: f64,
111    /// Estimated power consumption (Watts)
112    pub estimated_power_watts: f64,
113    /// Enable early stopping when budget reached
114    pub early_stopping: bool,
115    /// Warning threshold (fraction of budget)
116    pub warning_threshold: f64,
117}
118
119impl Default for EnergyAwareConfig {
120    fn default() -> Self {
121        Self {
122            energy_budget_kwh: 10.0,      // 10 kWh default budget
123            estimated_power_watts: 250.0, // Typical GPU power
124            early_stopping: true,
125            warning_threshold: 0.9,
126        }
127    }
128}
129
130/// Energy-Aware Optimizer
131///
132/// Tracks energy consumption and adapts training based on energy budget.
133/// Implements early stopping and energy-efficient parameter updates.
134pub struct EnergyAwareOptimizer<O: Optimizer> {
135    /// Base optimizer
136    base_optimizer: O,
137    /// Configuration
138    config: EnergyAwareConfig,
139    /// Energy metrics
140    metrics: EnergyMetrics,
141    /// Last step timestamp
142    last_step_time: Option<Instant>,
143    /// Budget exceeded flag
144    budget_exceeded: bool,
145}
146
147impl<O: Optimizer> EnergyAwareOptimizer<O> {
148    /// Create a new energy-aware optimizer
149    pub fn new(base_optimizer: O, config: EnergyAwareConfig) -> Self {
150        Self {
151            base_optimizer,
152            config,
153            metrics: EnergyMetrics::default(),
154            last_step_time: None,
155            budget_exceeded: false,
156        }
157    }
158
159    /// Create with default configuration
160    pub fn with_defaults(base_optimizer: O) -> Self {
161        Self::new(base_optimizer, EnergyAwareConfig::default())
162    }
163
164    /// Get current energy metrics
165    pub fn get_metrics(&self) -> &EnergyMetrics {
166        &self.metrics
167    }
168
169    /// Check if energy budget exceeded
170    pub fn is_budget_exceeded(&self) -> bool {
171        self.budget_exceeded
172    }
173
174    /// Update energy metrics
175    fn update_energy_metrics(&mut self, step_duration: Duration) {
176        self.metrics.num_steps += 1;
177        self.metrics.total_time += step_duration;
178
179        // Estimate energy consumption
180        // Energy (Joules) = Power (Watts) * Time (seconds)
181        let step_energy_joules = self.config.estimated_power_watts * step_duration.as_secs_f64();
182        let step_energy_kwh = step_energy_joules / 3_600_000.0; // Convert J to kWh
183
184        self.metrics.total_energy_kwh += step_energy_kwh;
185        self.metrics.energy_per_step = step_energy_joules;
186
187        // Update average power
188        if self.metrics.total_time.as_secs_f64() > 0.0 {
189            self.metrics.avg_power_watts = (self.metrics.total_energy_kwh * 3_600_000.0)
190                / self.metrics.total_time.as_secs_f64();
191        }
192
193        // Update peak power (estimated based on step duration)
194        let current_power = step_energy_joules / step_duration.as_secs_f64();
195        if current_power > self.metrics.peak_power_watts {
196            self.metrics.peak_power_watts = current_power;
197        }
198
199        // Check budget
200        if self.metrics.total_energy_kwh >= self.config.energy_budget_kwh {
201            self.budget_exceeded = true;
202        }
203
204        // Warning if approaching budget
205        let budget_fraction = self.metrics.total_energy_kwh / self.config.energy_budget_kwh;
206        if budget_fraction >= self.config.warning_threshold {
207            log::warn!(
208                "Energy budget {}% consumed: {:.3} / {:.3} kWh",
209                (budget_fraction * 100.0) as u32,
210                self.metrics.total_energy_kwh,
211                self.config.energy_budget_kwh
212            );
213        }
214    }
215
216    /// Get energy efficiency (steps per kWh)
217    pub fn get_efficiency(&self) -> f64 {
218        if self.metrics.total_energy_kwh > 0.0 {
219            self.metrics.num_steps as f64 / self.metrics.total_energy_kwh
220        } else {
221            0.0
222        }
223    }
224
225    /// Get estimated remaining steps
226    pub fn get_remaining_steps(&self) -> usize {
227        let remaining_energy = self.config.energy_budget_kwh - self.metrics.total_energy_kwh;
228        if remaining_energy > 0.0 && self.metrics.energy_per_step > 0.0 {
229            ((remaining_energy * 3_600_000.0) / self.metrics.energy_per_step) as usize
230        } else {
231            0
232        }
233    }
234}
235
236impl<O: Optimizer> Optimizer for EnergyAwareOptimizer<O> {
237    fn step(&mut self) -> OptimizerResult<()> {
238        // Check budget before step
239        if self.budget_exceeded && self.config.early_stopping {
240            return Err(OptimizerError::ConfigError(
241                "Energy budget exceeded".to_string(),
242            ));
243        }
244
245        let start = Instant::now();
246
247        // Perform optimization step
248        self.base_optimizer.step()?;
249
250        let step_duration = start.elapsed();
251        self.update_energy_metrics(step_duration);
252        self.last_step_time = Some(Instant::now());
253
254        Ok(())
255    }
256
257    fn zero_grad(&mut self) {
258        self.base_optimizer.zero_grad();
259    }
260
261    fn get_lr(&self) -> Vec<f32> {
262        self.base_optimizer.get_lr()
263    }
264
265    fn set_lr(&mut self, lr: f32) {
266        self.base_optimizer.set_lr(lr);
267    }
268
269    fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
270        self.base_optimizer.add_param_group(params, options);
271    }
272
273    fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
274        self.base_optimizer.parameters()
275    }
276
277    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
278        let mut state = self.base_optimizer.state_dict()?;
279        state.optimizer_type = format!("EnergyAware({})", state.optimizer_type);
280        state.global_state.insert(
281            "total_energy_kwh".to_string(),
282            self.metrics.total_energy_kwh as f32,
283        );
284        state
285            .global_state
286            .insert("num_steps".to_string(), self.metrics.num_steps as f32);
287        Ok(state)
288    }
289
290    fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
291        self.base_optimizer.load_state_dict(state)
292    }
293}
294
295// ============================================================================
296// Carbon-Conscious Optimizer
297// ============================================================================
298
299/// Carbon-conscious configuration
300#[derive(Debug, Clone)]
301pub struct CarbonConsciousConfig {
302    /// Carbon budget (gCO2)
303    pub carbon_budget_gco2: f64,
304    /// Estimated power consumption (Watts)
305    pub estimated_power_watts: f64,
306    /// Enable adaptive scheduling
307    pub adaptive_scheduling: bool,
308    /// Carbon intensity threshold for pausing (gCO2/kWh)
309    pub intensity_threshold: f64,
310}
311
312impl Default for CarbonConsciousConfig {
313    fn default() -> Self {
314        Self {
315            carbon_budget_gco2: 5000.0, // 5 kg CO2
316            estimated_power_watts: 250.0,
317            adaptive_scheduling: false,
318            intensity_threshold: 600.0, // Pause if > 600 gCO2/kWh
319        }
320    }
321}
322
323/// Carbon-Conscious Optimizer
324///
325/// Tracks carbon emissions during training and implements carbon-aware scheduling.
326/// Can pause training during high-carbon periods if adaptive scheduling is enabled.
327pub struct CarbonConsciousOptimizer<O: Optimizer> {
328    /// Base optimizer
329    base_optimizer: O,
330    /// Configuration
331    config: CarbonConsciousConfig,
332    /// Current carbon intensity
333    carbon_intensity: CarbonIntensity,
334    /// Total carbon emitted (gCO2)
335    total_carbon_gco2: f64,
336    /// Energy metrics
337    energy_metrics: EnergyMetrics,
338    /// Last step timestamp
339    last_step_time: Option<Instant>,
340}
341
342impl<O: Optimizer> CarbonConsciousOptimizer<O> {
343    /// Create a new carbon-conscious optimizer
344    pub fn new(
345        base_optimizer: O,
346        config: CarbonConsciousConfig,
347        carbon_intensity: CarbonIntensity,
348    ) -> Self {
349        Self {
350            base_optimizer,
351            config,
352            carbon_intensity,
353            total_carbon_gco2: 0.0,
354            energy_metrics: EnergyMetrics::default(),
355            last_step_time: None,
356        }
357    }
358
359    /// Create with default configuration
360    pub fn with_defaults(base_optimizer: O) -> Self {
361        Self::new(
362            base_optimizer,
363            CarbonConsciousConfig::default(),
364            CarbonIntensity::default(),
365        )
366    }
367
368    /// Update carbon intensity (e.g., from grid API)
369    pub fn update_carbon_intensity(&mut self, intensity: CarbonIntensity) {
370        self.carbon_intensity = intensity;
371    }
372
373    /// Get total carbon emissions
374    pub fn get_total_carbon(&self) -> f64 {
375        self.total_carbon_gco2
376    }
377
378    /// Get carbon efficiency (steps per kg CO2)
379    pub fn get_carbon_efficiency(&self) -> f64 {
380        if self.total_carbon_gco2 > 0.0 {
381            self.energy_metrics.num_steps as f64 / (self.total_carbon_gco2 / 1000.0)
382        } else {
383            0.0
384        }
385    }
386
387    /// Check if current carbon intensity is acceptable
388    fn should_proceed(&self) -> bool {
389        !self.config.adaptive_scheduling
390            || self.carbon_intensity.intensity <= self.config.intensity_threshold
391    }
392
393    /// Update carbon emissions
394    fn update_carbon_metrics(&mut self, step_duration: Duration) {
395        self.energy_metrics.num_steps += 1;
396        self.energy_metrics.total_time += step_duration;
397
398        // Calculate energy consumed
399        let step_energy_joules = self.config.estimated_power_watts * step_duration.as_secs_f64();
400        let step_energy_kwh = step_energy_joules / 3_600_000.0;
401
402        self.energy_metrics.total_energy_kwh += step_energy_kwh;
403
404        // Calculate carbon emissions: CO2 = Energy (kWh) × Carbon Intensity (gCO2/kWh)
405        let step_carbon_gco2 = step_energy_kwh * self.carbon_intensity.intensity;
406        self.total_carbon_gco2 += step_carbon_gco2;
407
408        // Log if approaching budget
409        let carbon_fraction = self.total_carbon_gco2 / self.config.carbon_budget_gco2;
410        if carbon_fraction >= 0.9 {
411            log::warn!(
412                "Carbon budget {}% consumed: {:.1} / {:.1} g CO2",
413                (carbon_fraction * 100.0) as u32,
414                self.total_carbon_gco2,
415                self.config.carbon_budget_gco2
416            );
417        }
418    }
419}
420
421impl<O: Optimizer> Optimizer for CarbonConsciousOptimizer<O> {
422    fn step(&mut self) -> OptimizerResult<()> {
423        // Check if we should proceed based on carbon intensity
424        if !self.should_proceed() {
425            return Err(OptimizerError::ConfigError(format!(
426                "Carbon intensity too high: {} gCO2/kWh (threshold: {})",
427                self.carbon_intensity.intensity, self.config.intensity_threshold
428            )));
429        }
430
431        // Check carbon budget
432        if self.total_carbon_gco2 >= self.config.carbon_budget_gco2 {
433            return Err(OptimizerError::ConfigError(
434                "Carbon budget exceeded".to_string(),
435            ));
436        }
437
438        let start = Instant::now();
439
440        // Perform optimization step
441        self.base_optimizer.step()?;
442
443        let step_duration = start.elapsed();
444        self.update_carbon_metrics(step_duration);
445        self.last_step_time = Some(Instant::now());
446
447        Ok(())
448    }
449
450    fn zero_grad(&mut self) {
451        self.base_optimizer.zero_grad();
452    }
453
454    fn get_lr(&self) -> Vec<f32> {
455        self.base_optimizer.get_lr()
456    }
457
458    fn set_lr(&mut self, lr: f32) {
459        self.base_optimizer.set_lr(lr);
460    }
461
462    fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
463        self.base_optimizer.add_param_group(params, options);
464    }
465
466    fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
467        self.base_optimizer.parameters()
468    }
469
470    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
471        let mut state = self.base_optimizer.state_dict()?;
472        state.optimizer_type = format!("CarbonConscious({})", state.optimizer_type);
473        state.global_state.insert(
474            "total_carbon_gco2".to_string(),
475            self.total_carbon_gco2 as f32,
476        );
477        state.global_state.insert(
478            "carbon_intensity".to_string(),
479            self.carbon_intensity.intensity as f32,
480        );
481        Ok(state)
482    }
483
484    fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
485        self.base_optimizer.load_state_dict(state)
486    }
487}
488
489// ============================================================================
490// Power-Capped Optimizer
491// ============================================================================
492
493/// Power-capped configuration
494#[derive(Debug, Clone)]
495pub struct PowerCappedConfig {
496    /// Maximum power consumption (Watts)
497    pub power_cap_watts: f64,
498    /// Target average power (Watts)
499    pub target_power_watts: f64,
500    /// Enable dynamic adjustment
501    pub dynamic_adjustment: bool,
502    /// Learning rate adjustment factor
503    pub lr_adjustment_factor: f32,
504}
505
506impl Default for PowerCappedConfig {
507    fn default() -> Self {
508        Self {
509            power_cap_watts: 300.0,
510            target_power_watts: 250.0,
511            dynamic_adjustment: true,
512            lr_adjustment_factor: 0.9,
513        }
514    }
515}
516
517/// Power-Capped Optimizer
518///
519/// Enforces power consumption limits by adapting learning rates and
520/// implementing power-aware parameter updates.
521pub struct PowerCappedOptimizer<O: Optimizer> {
522    /// Base optimizer
523    base_optimizer: O,
524    /// Configuration
525    config: PowerCappedConfig,
526    /// Current power estimate (Watts)
527    current_power_watts: f64,
528    /// Power history (exponential moving average)
529    power_ema: f64,
530    /// Adjustment count
531    adjustment_count: usize,
532}
533
534impl<O: Optimizer> PowerCappedOptimizer<O> {
535    /// Create a new power-capped optimizer
536    pub fn new(base_optimizer: O, config: PowerCappedConfig) -> Self {
537        let target_power = config.target_power_watts;
538        Self {
539            base_optimizer,
540            config,
541            current_power_watts: target_power,
542            power_ema: target_power,
543            adjustment_count: 0,
544        }
545    }
546
547    /// Create with default configuration
548    pub fn with_defaults(base_optimizer: O) -> Self {
549        Self::new(base_optimizer, PowerCappedConfig::default())
550    }
551
552    /// Get current power estimate
553    pub fn get_current_power(&self) -> f64 {
554        self.current_power_watts
555    }
556
557    /// Update power estimate based on step duration
558    fn update_power_estimate(&mut self, step_duration: Duration) {
559        // Simple model: faster steps = higher power
560        let baseline_duration = 0.1; // 100ms baseline
561        let duration_ratio = baseline_duration / step_duration.as_secs_f64().max(0.001);
562
563        self.current_power_watts = self.config.target_power_watts * duration_ratio;
564
565        // Update EMA
566        let alpha = 0.1;
567        self.power_ema = alpha * self.current_power_watts + (1.0 - alpha) * self.power_ema;
568    }
569
570    /// Adjust learning rate based on power
571    fn adjust_learning_rate(&mut self) {
572        if !self.config.dynamic_adjustment {
573            return;
574        }
575
576        if self.power_ema > self.config.power_cap_watts {
577            // Reduce learning rate to lower power
578            let current_lr = self.base_optimizer.get_lr();
579            let new_lr = current_lr[0] * self.config.lr_adjustment_factor;
580            self.base_optimizer.set_lr(new_lr);
581            self.adjustment_count += 1;
582
583            log::info!(
584                "Power cap exceeded ({:.1}W > {:.1}W), reducing LR to {:.6}",
585                self.power_ema,
586                self.config.power_cap_watts,
587                new_lr
588            );
589        }
590    }
591}
592
593impl<O: Optimizer> Optimizer for PowerCappedOptimizer<O> {
594    fn step(&mut self) -> OptimizerResult<()> {
595        let start = Instant::now();
596
597        // Perform optimization step
598        self.base_optimizer.step()?;
599
600        let step_duration = start.elapsed();
601        self.update_power_estimate(step_duration);
602        self.adjust_learning_rate();
603
604        Ok(())
605    }
606
607    fn zero_grad(&mut self) {
608        self.base_optimizer.zero_grad();
609    }
610
611    fn get_lr(&self) -> Vec<f32> {
612        self.base_optimizer.get_lr()
613    }
614
615    fn set_lr(&mut self, lr: f32) {
616        self.base_optimizer.set_lr(lr);
617    }
618
619    fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
620        self.base_optimizer.add_param_group(params, options);
621    }
622
623    fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
624        self.base_optimizer.parameters()
625    }
626
627    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
628        let mut state = self.base_optimizer.state_dict()?;
629        state.optimizer_type = format!("PowerCapped({})", state.optimizer_type);
630        state.global_state.insert(
631            "current_power_watts".to_string(),
632            self.current_power_watts as f32,
633        );
634        state
635            .global_state
636            .insert("adjustment_count".to_string(), self.adjustment_count as f32);
637        Ok(state)
638    }
639
640    fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
641        self.base_optimizer.load_state_dict(state)
642    }
643}
644
645// ============================================================================
646// Tests
647// ============================================================================
648
649#[cfg(test)]
650mod tests {
651    use super::*;
652    use crate::sgd::SGD;
653    use torsh_tensor::creation::randn;
654
655    #[test]
656    fn test_energy_aware_config_default() {
657        let config = EnergyAwareConfig::default();
658        assert_eq!(config.energy_budget_kwh, 10.0);
659        assert_eq!(config.estimated_power_watts, 250.0);
660        assert!(config.early_stopping);
661    }
662
663    #[test]
664    fn test_energy_aware_optimizer_creation() -> OptimizerResult<()> {
665        let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
666        let base = SGD::new(vec![param], 0.01, None, None, None, false);
667
668        let optimizer = EnergyAwareOptimizer::with_defaults(base);
669        assert_eq!(optimizer.metrics.num_steps, 0);
670        assert!(!optimizer.is_budget_exceeded());
671        Ok(())
672    }
673
674    #[test]
675    fn test_energy_metrics() -> OptimizerResult<()> {
676        let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
677        let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
678
679        let mut optimizer = EnergyAwareOptimizer::with_defaults(base);
680
681        // Set gradient and step
682        {
683            let mut p = param.write();
684            let grad = randn::<f32>(&[5, 5])?;
685            p.set_grad(Some(grad));
686        }
687
688        optimizer.step()?;
689
690        let metrics = optimizer.get_metrics();
691        assert_eq!(metrics.num_steps, 1);
692        assert!(metrics.total_energy_kwh > 0.0);
693        Ok(())
694    }
695
696    #[test]
697    fn test_carbon_conscious_config_default() {
698        let config = CarbonConsciousConfig::default();
699        assert_eq!(config.carbon_budget_gco2, 5000.0);
700        assert_eq!(config.estimated_power_watts, 250.0);
701    }
702
703    #[test]
704    fn test_carbon_conscious_optimizer_creation() -> OptimizerResult<()> {
705        let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
706        let base = SGD::new(vec![param], 0.01, None, None, None, false);
707
708        let optimizer = CarbonConsciousOptimizer::with_defaults(base);
709        assert_eq!(optimizer.get_total_carbon(), 0.0);
710        Ok(())
711    }
712
713    #[test]
714    fn test_carbon_intensity_update() -> OptimizerResult<()> {
715        let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
716        let base = SGD::new(vec![param], 0.01, None, None, None, false);
717
718        let mut optimizer = CarbonConsciousOptimizer::with_defaults(base);
719
720        let new_intensity = CarbonIntensity {
721            intensity: 300.0,
722            region: "test".to_string(),
723        };
724
725        optimizer.update_carbon_intensity(new_intensity.clone());
726        assert_eq!(optimizer.carbon_intensity.intensity, 300.0);
727        Ok(())
728    }
729
730    #[test]
731    fn test_power_capped_config_default() {
732        let config = PowerCappedConfig::default();
733        assert_eq!(config.power_cap_watts, 300.0);
734        assert_eq!(config.target_power_watts, 250.0);
735        assert!(config.dynamic_adjustment);
736    }
737
738    #[test]
739    fn test_power_capped_optimizer_creation() -> OptimizerResult<()> {
740        let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
741        let base = SGD::new(vec![param], 0.01, None, None, None, false);
742
743        let optimizer = PowerCappedOptimizer::with_defaults(base);
744        assert!(optimizer.get_current_power() > 0.0);
745        Ok(())
746    }
747
748    #[test]
749    fn test_power_estimate_update() -> OptimizerResult<()> {
750        let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
751        let base = SGD::new(vec![param.clone()], 0.01, None, None, None, false);
752
753        let mut optimizer = PowerCappedOptimizer::with_defaults(base);
754
755        // Set gradient and step
756        {
757            let mut p = param.write();
758            let grad = randn::<f32>(&[5, 5])?;
759            p.set_grad(Some(grad));
760        }
761
762        let initial_power = optimizer.get_current_power();
763        optimizer.step()?;
764        let updated_power = optimizer.get_current_power();
765
766        // Power should be updated (may increase or decrease)
767        assert!(updated_power > 0.0);
768        assert_ne!(initial_power, updated_power);
769        Ok(())
770    }
771}