1use 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#[derive(Debug, Clone)]
56pub struct EnergyMetrics {
57 pub total_energy_kwh: f64,
59 pub avg_power_watts: f64,
61 pub peak_power_watts: f64,
63 pub num_steps: usize,
65 pub total_time: Duration,
67 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#[derive(Debug, Clone)]
86pub struct CarbonIntensity {
87 pub intensity: f64,
89 pub region: String,
91}
92
93impl Default for CarbonIntensity {
94 fn default() -> Self {
95 Self {
96 intensity: 475.0, region: "global".to_string(),
98 }
99 }
100}
101
102#[derive(Debug, Clone)]
108pub struct EnergyAwareConfig {
109 pub energy_budget_kwh: f64,
111 pub estimated_power_watts: f64,
113 pub early_stopping: bool,
115 pub warning_threshold: f64,
117}
118
119impl Default for EnergyAwareConfig {
120 fn default() -> Self {
121 Self {
122 energy_budget_kwh: 10.0, estimated_power_watts: 250.0, early_stopping: true,
125 warning_threshold: 0.9,
126 }
127 }
128}
129
130pub struct EnergyAwareOptimizer<O: Optimizer> {
135 base_optimizer: O,
137 config: EnergyAwareConfig,
139 metrics: EnergyMetrics,
141 last_step_time: Option<Instant>,
143 budget_exceeded: bool,
145}
146
147impl<O: Optimizer> EnergyAwareOptimizer<O> {
148 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 pub fn with_defaults(base_optimizer: O) -> Self {
161 Self::new(base_optimizer, EnergyAwareConfig::default())
162 }
163
164 pub fn get_metrics(&self) -> &EnergyMetrics {
166 &self.metrics
167 }
168
169 pub fn is_budget_exceeded(&self) -> bool {
171 self.budget_exceeded
172 }
173
174 fn update_energy_metrics(&mut self, step_duration: Duration) {
176 self.metrics.num_steps += 1;
177 self.metrics.total_time += step_duration;
178
179 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; self.metrics.total_energy_kwh += step_energy_kwh;
185 self.metrics.energy_per_step = step_energy_joules;
186
187 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 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 if self.metrics.total_energy_kwh >= self.config.energy_budget_kwh {
201 self.budget_exceeded = true;
202 }
203
204 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 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 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 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 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#[derive(Debug, Clone)]
301pub struct CarbonConsciousConfig {
302 pub carbon_budget_gco2: f64,
304 pub estimated_power_watts: f64,
306 pub adaptive_scheduling: bool,
308 pub intensity_threshold: f64,
310}
311
312impl Default for CarbonConsciousConfig {
313 fn default() -> Self {
314 Self {
315 carbon_budget_gco2: 5000.0, estimated_power_watts: 250.0,
317 adaptive_scheduling: false,
318 intensity_threshold: 600.0, }
320 }
321}
322
323pub struct CarbonConsciousOptimizer<O: Optimizer> {
328 base_optimizer: O,
330 config: CarbonConsciousConfig,
332 carbon_intensity: CarbonIntensity,
334 total_carbon_gco2: f64,
336 energy_metrics: EnergyMetrics,
338 last_step_time: Option<Instant>,
340}
341
342impl<O: Optimizer> CarbonConsciousOptimizer<O> {
343 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 pub fn with_defaults(base_optimizer: O) -> Self {
361 Self::new(
362 base_optimizer,
363 CarbonConsciousConfig::default(),
364 CarbonIntensity::default(),
365 )
366 }
367
368 pub fn update_carbon_intensity(&mut self, intensity: CarbonIntensity) {
370 self.carbon_intensity = intensity;
371 }
372
373 pub fn get_total_carbon(&self) -> f64 {
375 self.total_carbon_gco2
376 }
377
378 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 fn should_proceed(&self) -> bool {
389 !self.config.adaptive_scheduling
390 || self.carbon_intensity.intensity <= self.config.intensity_threshold
391 }
392
393 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 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 let step_carbon_gco2 = step_energy_kwh * self.carbon_intensity.intensity;
406 self.total_carbon_gco2 += step_carbon_gco2;
407
408 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 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 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 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#[derive(Debug, Clone)]
495pub struct PowerCappedConfig {
496 pub power_cap_watts: f64,
498 pub target_power_watts: f64,
500 pub dynamic_adjustment: bool,
502 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
517pub struct PowerCappedOptimizer<O: Optimizer> {
522 base_optimizer: O,
524 config: PowerCappedConfig,
526 current_power_watts: f64,
528 power_ema: f64,
530 adjustment_count: usize,
532}
533
534impl<O: Optimizer> PowerCappedOptimizer<O> {
535 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 pub fn with_defaults(base_optimizer: O) -> Self {
549 Self::new(base_optimizer, PowerCappedConfig::default())
550 }
551
552 pub fn get_current_power(&self) -> f64 {
554 self.current_power_watts
555 }
556
557 fn update_power_estimate(&mut self, step_duration: Duration) {
559 let baseline_duration = 0.1; 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 let alpha = 0.1;
567 self.power_ema = alpha * self.current_power_watts + (1.0 - alpha) * self.power_ema;
568 }
569
570 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 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 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#[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 {
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 {
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 assert!(updated_power > 0.0);
768 assert_ne!(initial_power, updated_power);
769 Ok(())
770 }
771}