1use crate::common::StateMemoryStats;
35use std::collections::HashMap;
36use trustformers_core::errors::{Result, TrustformersError};
37use trustformers_core::tensor::Tensor;
38use trustformers_core::traits::Optimizer;
39
40pub trait StatefulOptimizer: Optimizer {
45 type Config: Clone + Send + Sync;
47
48 type State: Send + Sync;
50
51 fn config(&self) -> &Self::Config;
53
54 fn state(&self) -> &Self::State;
56
57 fn state_mut(&mut self) -> &mut Self::State;
59
60 fn state_dict(&self) -> Result<HashMap<String, Tensor>>;
62
63 fn load_state_dict(&mut self, state: HashMap<String, Tensor>) -> Result<()>;
65
66 fn memory_usage(&self) -> StateMemoryStats;
68
69 fn reset_state(&mut self);
71
72 fn num_parameters(&self) -> usize;
74
75 fn save_state(&self, path: &std::path::Path) -> Result<()> {
85 let state = self.state_dict()?;
86 let encoded = encode_state_dict(&state)?;
87 std::fs::write(path, encoded).map_err(|error| {
88 TrustformersError::io_error(format!(
89 "failed to write optimizer state to {}: {error}",
90 path.display()
91 ))
92 })
93 }
94
95 fn load_state(&mut self, path: &std::path::Path) -> Result<()> {
102 let bytes = std::fs::read(path).map_err(|error| {
103 TrustformersError::io_error(format!(
104 "failed to read optimizer state from {}: {error}",
105 path.display()
106 ))
107 })?;
108 let state = decode_state_dict(&bytes)?;
109 self.load_state_dict(state)
110 }
111}
112
113type WireStateDict = Vec<(String, Vec<usize>, Vec<f32>)>;
118
119pub fn encode_state_dict(state: &HashMap<String, Tensor>) -> Result<Vec<u8>> {
125 let mut wire: WireStateDict = Vec::with_capacity(state.len());
126 for (name, tensor) in state {
127 wire.push((name.clone(), tensor.shape().to_vec(), tensor.data_f32()?));
128 }
129 wire.sort_by(|a, b| a.0.cmp(&b.0));
131
132 oxicode::serde::encode_to_vec(&wire, oxicode::config::standard()).map_err(|error| {
133 TrustformersError::invalid_state(format!("failed to encode optimizer state: {error}"))
134 })
135}
136
137pub fn decode_state_dict(bytes: &[u8]) -> Result<HashMap<String, Tensor>> {
144 let (wire, _): (WireStateDict, usize) =
145 oxicode::serde::decode_from_slice(bytes, oxicode::config::standard()).map_err(|error| {
146 TrustformersError::invalid_state(format!("failed to decode optimizer state: {error}"))
147 })?;
148
149 let mut state = HashMap::with_capacity(wire.len());
150 for (name, shape, values) in wire {
151 let expected: usize = shape.iter().product();
152 if values.len() != expected {
153 return Err(TrustformersError::invalid_state(format!(
154 "optimizer state entry '{name}' has {} values but shape {shape:?} needs {expected}",
155 values.len()
156 )));
157 }
158 state.insert(name, Tensor::from_vec(values, &shape)?);
159 }
160 Ok(state)
161}
162
163pub trait MomentumOptimizer: StatefulOptimizer {
167 fn momentum_coeff(&self) -> f32;
169
170 fn set_momentum_coeff(&mut self, coeff: f32);
172
173 fn momentum_buffers(&self) -> &HashMap<String, Vec<f32>>;
175
176 fn clear_momentum(&mut self);
178}
179
180pub trait AdaptiveMomentumOptimizer: MomentumOptimizer {
184 fn variance_coeff(&self) -> f32;
186
187 fn set_variance_coeff(&mut self, coeff: f32);
189
190 fn epsilon(&self) -> f32;
192
193 fn set_epsilon(&mut self, eps: f32);
195
196 fn variance_buffers(&self) -> &HashMap<String, Vec<f32>>;
198
199 fn clear_variance(&mut self);
201
202 fn apply_bias_correction(&self, momentum: f32, variance: f32, step: usize) -> (f32, f32);
204}
205
206pub trait ClassicalMomentumOptimizer: MomentumOptimizer {
208 fn dampening(&self) -> f32;
210
211 fn set_dampening(&mut self, dampening: f32);
213
214 fn nesterov(&self) -> bool;
216
217 fn set_nesterov(&mut self, nesterov: bool);
219}
220
221pub trait SecondOrderOptimizer: StatefulOptimizer {
225 type CurvatureInfo;
227
228 fn update_curvature(&mut self, gradients: &[Tensor]) -> Result<()>;
230
231 fn curvature_info(&self) -> &Self::CurvatureInfo;
233
234 fn apply_inverse_hessian(&self, gradient: &Tensor) -> Result<Tensor>;
236
237 fn history_size(&self) -> usize;
239}
240
241pub trait DistributedOptimizer: Optimizer {
245 type Communicator;
247
248 fn all_reduce_gradients(&mut self, gradients: &mut [Tensor]) -> Result<()>;
250
251 fn broadcast_parameters(&mut self, parameters: &mut [Tensor]) -> Result<()>;
253
254 fn rank(&self) -> usize;
256
257 fn world_size(&self) -> usize;
259
260 fn sync_state(&mut self) -> Result<()>;
262}
263
264pub trait GradientCompressionOptimizer: DistributedOptimizer {
266 type CompressionMethod;
268
269 fn compress_gradients(&self, gradients: &[Tensor]) -> Result<Vec<u8>>;
271
272 fn decompress_gradients(&self, data: &[u8]) -> Result<Vec<Tensor>>;
274
275 fn compression_ratio(&self) -> f32;
277
278 fn set_compression_config(&mut self, config: Self::CompressionMethod);
280}
281
282pub trait FederatedOptimizer: DistributedOptimizer {
284 type ClientInfo;
286
287 fn aggregate_updates(
289 &mut self,
290 updates: &[Tensor],
291 clients: &[Self::ClientInfo],
292 ) -> Result<Tensor>;
293
294 fn select_clients(
296 &self,
297 available_clients: &[Self::ClientInfo],
298 num_clients: usize,
299 ) -> Vec<usize>;
300
301 fn apply_differential_privacy(&mut self, update: &mut Tensor) -> Result<()>;
303}
304
305pub trait AsyncOptimizer: DistributedOptimizer {
307 fn apply_delayed_gradients(&mut self, gradients: &[Tensor], staleness: usize) -> Result<()>;
309
310 fn max_staleness(&self) -> usize;
312
313 fn set_staleness_compensation(&mut self, method: StalenessCompensation);
315}
316
317#[derive(Debug, Clone, Copy)]
319pub enum StalenessCompensation {
320 None,
322 Linear,
324 Exponential,
326 Polynomial(f32),
328}
329
330pub trait HardwareOptimizer: Optimizer {
332 type HardwareTarget;
334
335 fn optimize_for_hardware(&mut self, target: Self::HardwareTarget) -> Result<()>;
337
338 fn hardware_utilization(&self) -> HardwareStats;
340
341 fn is_hardware_compatible(&self) -> bool;
343}
344
345#[derive(Debug, Clone)]
347pub struct HardwareStats {
348 pub memory_bandwidth_utilization: f32,
350 pub compute_utilization: f32,
352 pub cache_hit_rate: f32,
354 pub flops_per_second: f64,
356}
357
358pub trait SIMDOptimizer: HardwareOptimizer {
360 type SIMDType;
362
363 fn simd_available(&self) -> bool;
365
366 fn vector_width(&self) -> usize;
368
369 fn simd_update(&mut self, parameters: &mut [Tensor], gradients: &[Tensor]) -> Result<()>;
371}
372
373pub trait GPUOptimizer: HardwareOptimizer {
375 type ComputeCapability;
377
378 fn to_gpu(&mut self) -> Result<()>;
380
381 fn to_cpu(&mut self) -> Result<()>;
383
384 fn gpu_update(&mut self, parameters: &mut [Tensor], gradients: &[Tensor]) -> Result<()>;
386
387 fn gpu_memory_usage(&self) -> GPUMemoryStats;
389}
390
391#[derive(Debug, Clone)]
393pub struct GPUMemoryStats {
394 pub total_memory: usize,
396 pub used_memory: usize,
398 pub available_memory: usize,
400 pub optimizer_memory: usize,
402}
403
404pub trait EdgeOptimizer: HardwareOptimizer {
406 type PowerStats;
408
409 fn optimize_for_power(&mut self) -> Result<()>;
411
412 fn power_stats(&self) -> Self::PowerStats;
414
415 fn reduce_precision(&mut self, bits: u8) -> Result<()>;
417}
418
419pub trait MetaOptimizer: Optimizer {
421 type BaseOptimizer: Optimizer;
423
424 fn base_optimizer(&self) -> &Self::BaseOptimizer;
426
427 fn base_optimizer_mut(&mut self) -> &mut Self::BaseOptimizer;
429
430 fn apply_meta_strategy(
432 &mut self,
433 parameters: &mut [Tensor],
434 gradients: &[Tensor],
435 ) -> Result<()>;
436}
437
438pub trait LookaheadOptimizer: MetaOptimizer {
440 fn lookahead_alpha(&self) -> f32;
442
443 fn set_lookahead_alpha(&mut self, alpha: f32);
445
446 fn lookahead_k(&self) -> usize;
448
449 fn set_lookahead_k(&mut self, k: usize);
451
452 fn slow_weights(&self) -> &HashMap<String, Vec<f32>>;
454}
455
456pub trait ScheduledOptimizer: Optimizer {
458 type Scheduler;
460
461 fn scheduler(&self) -> &Self::Scheduler;
463
464 fn scheduler_mut(&mut self) -> &mut Self::Scheduler;
466
467 fn update_lr(&mut self) -> Result<()>;
469
470 fn current_lr(&self) -> f32;
472}
473
474pub trait CompositeOptimizer: Optimizer {
476 type Components;
478
479 fn components(&self) -> &Self::Components;
481
482 fn components_mut(&mut self) -> &mut Self::Components;
484
485 fn apply_composite_update(
487 &mut self,
488 parameters: &mut [Tensor],
489 gradients: &[Tensor],
490 ) -> Result<()>;
491
492 fn component_weights(&self) -> Vec<f32>;
494
495 fn set_component_weights(&mut self, weights: Vec<f32>) -> Result<()>;
497}
498
499pub trait OptimizerFactory {
501 type Optimizer: Optimizer;
503
504 type Config;
506
507 fn create(&self, config: Self::Config) -> Result<Self::Optimizer>;
509
510 fn available_variants(&self) -> Vec<&'static str>;
512
513 fn create_by_name(&self, name: &str) -> Result<Self::Optimizer>;
515}
516
517pub trait SerializableOptimizer: Optimizer {
519 fn serialize(&self) -> Result<Vec<u8>>;
521
522 fn deserialize(data: &[u8]) -> Result<Self>
524 where
525 Self: Sized;
526
527 fn version(&self) -> u32;
529}
530
531#[cfg(test)]
532mod tests {
533 use super::*;
534
535 #[test]
536 fn test_staleness_compensation() {
537 let compensation = StalenessCompensation::Linear;
538 assert!(
539 matches!(compensation, StalenessCompensation::Linear),
540 "Expected Linear staleness compensation"
541 );
542 }
543
544 #[test]
545 fn test_hardware_stats() {
546 let stats = HardwareStats {
547 memory_bandwidth_utilization: 0.8,
548 compute_utilization: 0.9,
549 cache_hit_rate: 0.95,
550 flops_per_second: 1e12,
551 };
552
553 assert_eq!(stats.memory_bandwidth_utilization, 0.8);
554 assert_eq!(stats.compute_utilization, 0.9);
555 assert_eq!(stats.cache_hit_rate, 0.95);
556 assert_eq!(stats.flops_per_second, 1e12);
557 }
558
559 #[test]
560 fn test_gpu_memory_stats() {
561 let stats = GPUMemoryStats {
562 total_memory: 16 * 1024 * 1024 * 1024, used_memory: 8 * 1024 * 1024 * 1024, available_memory: 8 * 1024 * 1024 * 1024, optimizer_memory: 1024 * 1024 * 1024, };
567
568 assert_eq!(stats.total_memory, 16 * 1024 * 1024 * 1024);
569 assert_eq!(
570 stats.used_memory + stats.available_memory,
571 stats.total_memory
572 );
573 assert!(stats.optimizer_memory <= stats.used_memory);
574 }
575
576 #[test]
577 fn test_staleness_compensation_none() {
578 let comp = StalenessCompensation::None;
579 assert!(matches!(comp, StalenessCompensation::None));
580 }
581
582 #[test]
583 fn test_staleness_compensation_exponential() {
584 let comp = StalenessCompensation::Exponential;
585 assert!(matches!(comp, StalenessCompensation::Exponential));
586 }
587
588 #[test]
589 fn test_staleness_compensation_polynomial() {
590 let comp = StalenessCompensation::Polynomial(2.0);
591 if let StalenessCompensation::Polynomial(degree) = comp {
592 assert_eq!(degree, 2.0);
593 } else {
594 panic!("Expected Polynomial variant");
595 }
596 }
597
598 #[test]
599 fn test_hardware_stats_all_zero() {
600 let stats = HardwareStats {
601 memory_bandwidth_utilization: 0.0,
602 compute_utilization: 0.0,
603 cache_hit_rate: 0.0,
604 flops_per_second: 0.0,
605 };
606 assert_eq!(stats.memory_bandwidth_utilization, 0.0);
607 assert_eq!(stats.flops_per_second, 0.0);
608 }
609
610 #[test]
611 fn test_hardware_stats_max_utilization() {
612 let stats = HardwareStats {
613 memory_bandwidth_utilization: 1.0,
614 compute_utilization: 1.0,
615 cache_hit_rate: 1.0,
616 flops_per_second: 1e15,
617 };
618 assert!(stats.memory_bandwidth_utilization <= 1.0);
619 assert!(stats.compute_utilization <= 1.0);
620 assert!(stats.cache_hit_rate <= 1.0);
621 }
622
623 #[test]
624 fn test_gpu_memory_stats_zero_usage() {
625 let stats = GPUMemoryStats {
626 total_memory: 16 * 1024 * 1024 * 1024,
627 used_memory: 0,
628 available_memory: 16 * 1024 * 1024 * 1024,
629 optimizer_memory: 0,
630 };
631 assert_eq!(stats.used_memory, 0);
632 assert_eq!(stats.total_memory, stats.available_memory);
633 }
634
635 #[test]
636 fn test_gpu_memory_stats_full_usage() {
637 let total = 8 * 1024 * 1024 * 1024_usize;
638 let stats = GPUMemoryStats {
639 total_memory: total,
640 used_memory: total,
641 available_memory: 0,
642 optimizer_memory: total / 4,
643 };
644 assert_eq!(stats.available_memory, 0);
645 assert!(stats.optimizer_memory <= stats.used_memory);
646 }
647
648 #[test]
649 fn test_gpu_memory_stats_optimizer_fraction() {
650 let total = 16 * 1024 * 1024 * 1024_usize;
651 let used = 12 * 1024 * 1024 * 1024_usize;
652 let optimizer = 3 * 1024 * 1024 * 1024_usize;
653 let stats = GPUMemoryStats {
654 total_memory: total,
655 used_memory: used,
656 available_memory: total - used,
657 optimizer_memory: optimizer,
658 };
659 assert_eq!(stats.available_memory, 4 * 1024 * 1024 * 1024);
660 assert!(stats.optimizer_memory < stats.used_memory);
661 }
662
663 #[test]
664 fn test_hardware_stats_clone() {
665 let stats = HardwareStats {
666 memory_bandwidth_utilization: 0.5,
667 compute_utilization: 0.7,
668 cache_hit_rate: 0.9,
669 flops_per_second: 5e11,
670 };
671 let cloned = stats.clone();
672 assert_eq!(cloned.memory_bandwidth_utilization, 0.5);
673 assert_eq!(cloned.compute_utilization, 0.7);
674 }
675
676 #[test]
677 fn test_gpu_memory_stats_clone() {
678 let stats = GPUMemoryStats {
679 total_memory: 1000,
680 used_memory: 500,
681 available_memory: 500,
682 optimizer_memory: 100,
683 };
684 let cloned = stats.clone();
685 assert_eq!(cloned.total_memory, 1000);
686 }
687
688 #[test]
689 fn test_staleness_compensation_copy() {
690 let comp = StalenessCompensation::Linear;
691 let copied = comp;
692 assert!(matches!(copied, StalenessCompensation::Linear));
693 }
694
695 #[test]
696 fn test_hardware_stats_realistic_gpu() {
697 let stats = HardwareStats {
698 memory_bandwidth_utilization: 0.75,
699 compute_utilization: 0.85,
700 cache_hit_rate: 0.92,
701 flops_per_second: 1.2e13,
702 };
703 assert!(
704 stats.memory_bandwidth_utilization > 0.0 && stats.memory_bandwidth_utilization <= 1.0
705 );
706 assert!(stats.compute_utilization > 0.0 && stats.compute_utilization <= 1.0);
707 assert!(stats.flops_per_second > 1e12);
708 }
709
710 #[test]
711 fn test_hardware_stats_edge_device() {
712 let stats = HardwareStats {
713 memory_bandwidth_utilization: 0.3,
714 compute_utilization: 0.4,
715 cache_hit_rate: 0.6,
716 flops_per_second: 1e9,
717 };
718 assert!(stats.flops_per_second < 1e10);
719 assert!(stats.compute_utilization < 0.5);
720 }
721
722 #[test]
723 fn test_gpu_memory_stats_consistency() {
724 let stats = GPUMemoryStats {
725 total_memory: 8 * 1024 * 1024 * 1024,
726 used_memory: 6 * 1024 * 1024 * 1024,
727 available_memory: 2 * 1024 * 1024 * 1024,
728 optimizer_memory: 2 * 1024 * 1024 * 1024,
729 };
730 assert_eq!(
731 stats.used_memory + stats.available_memory,
732 stats.total_memory
733 );
734 }
735
736 #[test]
737 fn test_staleness_polynomial_fractional() {
738 let comp = StalenessCompensation::Polynomial(0.5);
739 if let StalenessCompensation::Polynomial(degree) = comp {
740 assert!(degree > 0.0 && degree < 1.0);
741 }
742 }
743
744 #[test]
745 fn test_staleness_polynomial_high_degree() {
746 let comp = StalenessCompensation::Polynomial(10.0);
747 if let StalenessCompensation::Polynomial(degree) = comp {
748 assert!(degree > 5.0);
749 }
750 }
751
752 #[test]
753 fn test_hardware_stats_debug_format() {
754 let stats = HardwareStats {
755 memory_bandwidth_utilization: 0.5,
756 compute_utilization: 0.5,
757 cache_hit_rate: 0.5,
758 flops_per_second: 1.0,
759 };
760 let debug_str = format!("{:?}", stats);
761 assert!(debug_str.contains("HardwareStats"));
762 }
763
764 #[test]
765 fn test_gpu_memory_stats_debug_format() {
766 let stats = GPUMemoryStats {
767 total_memory: 100,
768 used_memory: 50,
769 available_memory: 50,
770 optimizer_memory: 10,
771 };
772 let debug_str = format!("{:?}", stats);
773 assert!(debug_str.contains("GPUMemoryStats"));
774 }
775}
776
777#[cfg(test)]
778mod state_persistence_tests {
779 use super::*;
780 use crate::adam::Adam;
781
782 #[test]
786 fn save_state_and_load_state_round_trip() {
787 let mut optimizer = Adam::new(0.01, (0.9, 0.999), 1e-8, 0.0);
788 let mut param = Tensor::from_vec(vec![1.0_f32, 2.0], &[2]).expect("tensor");
789 let grad = Tensor::from_vec(vec![0.5_f32, -0.5], &[2]).expect("grad");
790 optimizer.update_named("w", &mut param, &grad).expect("step 1");
791 Optimizer::step(&mut optimizer);
792 optimizer.update_named("w", &mut param, &grad).expect("step 2");
793
794 let path = std::env::temp_dir().join(format!(
795 "trustformers-optim-state-{}.bin",
796 std::process::id()
797 ));
798 optimizer.save_state(&path).expect("save_state");
799
800 let mut restored = Adam::new(0.5, (0.1, 0.1), 1e-2, 0.9);
801 restored.load_state(&path).expect("load_state");
802 let _ = std::fs::remove_file(&path);
803
804 let original = optimizer.state_dict().expect("state_dict");
805 let round_trip = restored.state_dict().expect("state_dict");
806 assert_eq!(original.len(), round_trip.len(), "every entry must survive");
807 for (key, tensor) in &original {
808 let other = round_trip.get(key).unwrap_or_else(|| panic!("missing '{key}'"));
809 assert_eq!(other.shape(), tensor.shape(), "shape of '{key}'");
810 assert_eq!(
811 other.data_f32().expect("data"),
812 tensor.data_f32().expect("data"),
813 "payload of '{key}'"
814 );
815 }
816 }
817
818 #[test]
820 fn decoding_rejects_corrupt_state() {
821 assert!(decode_state_dict(&[0xff, 0x00, 0x13, 0x37]).is_err());
822 }
823
824 #[test]
826 fn decoding_rejects_a_shape_payload_mismatch() {
827 let wire: Vec<(String, Vec<usize>, Vec<f32>)> =
828 vec![("w".to_string(), vec![4], vec![1.0, 2.0])];
829 let bytes =
830 oxicode::serde::encode_to_vec(&wire, oxicode::config::standard()).expect("encode");
831 assert!(decode_state_dict(&bytes).is_err());
832 }
833
834 #[test]
836 fn load_state_reports_a_missing_file() {
837 let mut optimizer = Adam::new(0.01, (0.9, 0.999), 1e-8, 0.0);
838 let path = std::env::temp_dir().join("trustformers-optim-definitely-absent.bin");
839 let _ = std::fs::remove_file(&path);
840 assert!(optimizer.load_state(&path).is_err());
841 }
842}