1#![allow(dead_code)]
23
24use crate::common::{OptimizerState, StateMemoryStats};
25use crate::traits::StatefulOptimizer;
26use serde::{Deserialize, Serialize};
27use std::collections::HashMap;
28use trustformers_core::errors::{Result, TrustformersError};
29use trustformers_core::tensor::Tensor;
30use trustformers_core::traits::Optimizer;
31
32#[derive(Debug, Clone, Serialize, Deserialize)]
34pub struct MicroAdamConfig {
35 pub learning_rate: f32,
37 pub beta1: f32,
39 pub beta2: f32,
41 pub epsilon: f32,
43 pub weight_decay: f32,
45 pub compression_ratio: f32,
47 pub min_block_size: usize,
49 pub adaptive_compression: bool,
51 pub compression_threshold: f32,
53 pub bias_correction: bool,
55 pub max_compression_error: f32,
57}
58
59impl Default for MicroAdamConfig {
60 fn default() -> Self {
61 Self {
62 learning_rate: 1e-3,
63 beta1: 0.9,
64 beta2: 0.999,
65 epsilon: 1e-8,
66 weight_decay: 0.01,
67 compression_ratio: 0.1,
68 min_block_size: 64,
69 adaptive_compression: true,
70 compression_threshold: 1e-6,
71 bias_correction: true,
72 max_compression_error: 1e-4,
73 }
74 }
75}
76
77#[derive(Debug, Clone)]
79struct CompressedGradient {
80 compressed_data: Vec<f32>,
82 indices: Vec<usize>,
84 scale_factor: f32,
86 original_size: usize,
88 compression_type: CompressionType,
90}
91
92#[derive(Debug, Clone, Copy)]
94enum CompressionType {
95 TopK,
97 Threshold,
99 BlockWise,
101 Adaptive,
103}
104
105impl CompressedGradient {
106 fn compress(gradient: &[f32], config: &MicroAdamConfig) -> Self {
108 let original_size = gradient.len();
109 let target_size = (original_size as f32 * config.compression_ratio) as usize;
110 let target_size = target_size.max(config.min_block_size.min(original_size));
111
112 let compression_type = if config.adaptive_compression {
113 Self::choose_adaptive_compression(gradient, config)
115 } else {
116 CompressionType::TopK
117 };
118
119 match compression_type {
120 CompressionType::TopK => Self::compress_topk(gradient, target_size),
121 CompressionType::Threshold => Self::compress_threshold(gradient, config),
122 CompressionType::BlockWise => Self::compress_blockwise(gradient, config),
123 CompressionType::Adaptive => Self::compress_adaptive(gradient, config),
124 }
125 }
126
127 fn choose_adaptive_compression(gradient: &[f32], config: &MicroAdamConfig) -> CompressionType {
129 let mean_abs = gradient.iter().map(|x| x.abs()).sum::<f32>() / gradient.len() as f32;
130 let sparsity = gradient.iter().filter(|&&x| x.abs() < config.compression_threshold).count()
131 as f32
132 / gradient.len() as f32;
133
134 if sparsity > 0.8 {
135 CompressionType::Threshold
136 } else if mean_abs > 1e-3 {
137 CompressionType::BlockWise
138 } else {
139 CompressionType::TopK
140 }
141 }
142
143 fn compress_topk(gradient: &[f32], k: usize) -> Self {
145 let mut indexed_values: Vec<(usize, f32)> =
146 gradient.iter().enumerate().map(|(i, &val)| (i, val.abs())).collect();
147
148 indexed_values.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
150
151 let k = k.min(indexed_values.len());
152 let indices: Vec<usize> = indexed_values[..k].iter().map(|(i, _)| *i).collect();
153 let compressed_data: Vec<f32> = indices.iter().map(|&i| gradient[i]).collect();
154
155 let max_val = compressed_data.iter().map(|x| x.abs()).fold(0.0f32, f32::max);
157 let scale_factor = if max_val > 0.0 { 1.0 / max_val } else { 1.0 };
158
159 Self {
160 compressed_data: compressed_data.iter().map(|x| x * scale_factor).collect(),
161 indices,
162 scale_factor: 1.0 / scale_factor,
163 original_size: gradient.len(),
164 compression_type: CompressionType::TopK,
165 }
166 }
167
168 fn compress_threshold(gradient: &[f32], config: &MicroAdamConfig) -> Self {
170 let threshold = config.compression_threshold;
171 let mut indices = Vec::new();
172 let mut compressed_data = Vec::new();
173
174 for (i, &val) in gradient.iter().enumerate() {
175 if val.abs() >= threshold {
176 indices.push(i);
177 compressed_data.push(val);
178 }
179 }
180
181 let max_val = compressed_data.iter().map(|x| x.abs()).fold(0.0f32, f32::max);
182 let scale_factor = if max_val > 0.0 { 1.0 / max_val } else { 1.0 };
183
184 Self {
185 compressed_data: compressed_data.iter().map(|x| x * scale_factor).collect(),
186 indices,
187 scale_factor: 1.0 / scale_factor,
188 original_size: gradient.len(),
189 compression_type: CompressionType::Threshold,
190 }
191 }
192
193 fn compress_blockwise(gradient: &[f32], config: &MicroAdamConfig) -> Self {
195 let block_size = config.min_block_size;
196 let num_blocks = gradient.len().div_ceil(block_size);
197 let target_elements_per_block =
198 ((block_size as f32 * config.compression_ratio) as usize).max(1);
199
200 let mut indices = Vec::new();
201 let mut compressed_data = Vec::new();
202
203 for block_idx in 0..num_blocks {
204 let start = block_idx * block_size;
205 let end = (start + block_size).min(gradient.len());
206 let block = &gradient[start..end];
207
208 let mut block_indexed: Vec<(usize, f32)> =
210 block.iter().enumerate().map(|(i, &val)| (start + i, val.abs())).collect();
211
212 block_indexed
213 .sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
214
215 let k = target_elements_per_block.min(block_indexed.len());
216 for i in 0..k {
217 let global_idx = block_indexed[i].0;
218 indices.push(global_idx);
219 compressed_data.push(gradient[global_idx]);
220 }
221 }
222
223 let max_val = compressed_data.iter().map(|x| x.abs()).fold(0.0f32, f32::max);
224 let scale_factor = if max_val > 0.0 { 1.0 / max_val } else { 1.0 };
225
226 Self {
227 compressed_data: compressed_data.iter().map(|x| x * scale_factor).collect(),
228 indices,
229 scale_factor: 1.0 / scale_factor,
230 original_size: gradient.len(),
231 compression_type: CompressionType::BlockWise,
232 }
233 }
234
235 fn compress_adaptive(gradient: &[f32], config: &MicroAdamConfig) -> Self {
237 let topk = Self::compress_topk(
239 gradient,
240 (gradient.len() as f32 * config.compression_ratio) as usize,
241 );
242 let threshold = Self::compress_threshold(gradient, config);
243 let blockwise = Self::compress_blockwise(gradient, config);
244
245 let topk_ratio = topk.compressed_data.len() as f32 / gradient.len() as f32;
247 let threshold_ratio = threshold.compressed_data.len() as f32 / gradient.len() as f32;
248 let blockwise_ratio = blockwise.compressed_data.len() as f32 / gradient.len() as f32;
249
250 if threshold_ratio <= config.compression_ratio && threshold_ratio < topk_ratio {
251 threshold
252 } else if blockwise_ratio <= config.compression_ratio && blockwise_ratio < topk_ratio {
253 blockwise
254 } else {
255 topk
256 }
257 }
258
259 fn decompress(&self) -> Vec<f32> {
261 let mut result = vec![0.0; self.original_size];
262 for (i, &idx) in self.indices.iter().enumerate() {
263 if idx < self.original_size && i < self.compressed_data.len() {
264 result[idx] = self.compressed_data[i] * self.scale_factor;
265 }
266 }
267 result
268 }
269
270 fn compression_ratio(&self) -> f32 {
272 self.compressed_data.len() as f32 / self.original_size as f32
273 }
274
275 fn compression_error(&self, original: &[f32]) -> f32 {
277 let decompressed = self.decompress();
278 let mut error_sum = 0.0;
279 let mut norm_sum = 0.0;
280
281 for (orig, decomp) in original.iter().zip(decompressed.iter()) {
282 error_sum += (orig - decomp).powi(2);
283 norm_sum += orig.powi(2);
284 }
285
286 if norm_sum > 0.0 {
287 (error_sum / norm_sum).sqrt()
288 } else {
289 0.0
290 }
291 }
292}
293
294#[derive(Debug)]
299pub struct MicroAdam {
300 config: MicroAdamConfig,
301 state: OptimizerState,
302 momentum: HashMap<String, CompressedGradient>,
304 variance: HashMap<String, CompressedGradient>,
306 compression_stats: CompressionStats,
308}
309
310#[derive(Debug, Default)]
312struct CompressionStats {
313 total_parameters: usize,
314 total_compressed_size: usize,
315 average_compression_ratio: f32,
316 average_compression_error: f32,
317 compression_method_usage: HashMap<String, usize>,
318}
319
320impl MicroAdam {
321 pub fn new() -> Self {
323 Self::with_config(MicroAdamConfig::default())
324 }
325
326 pub fn new_with_lr(learning_rate: f32) -> Self {
328 let config = MicroAdamConfig {
329 learning_rate,
330 ..Default::default()
331 };
332 Self::with_config(config)
333 }
334
335 pub fn for_large_models() -> Self {
337 let config = MicroAdamConfig {
338 learning_rate: 1e-4,
339 beta1: 0.9,
340 beta2: 0.999,
341 epsilon: 1e-8,
342 weight_decay: 0.01,
343 compression_ratio: 0.05, min_block_size: 128,
345 adaptive_compression: true,
346 compression_threshold: 1e-7,
347 bias_correction: true,
348 max_compression_error: 1e-5,
349 };
350 Self::with_config(config)
351 }
352
353 pub fn for_memory_constrained() -> Self {
355 let config = MicroAdamConfig {
356 learning_rate: 1e-3,
357 beta1: 0.9,
358 beta2: 0.999,
359 epsilon: 1e-8,
360 weight_decay: 0.01,
361 compression_ratio: 0.02, min_block_size: 32,
363 adaptive_compression: true,
364 compression_threshold: 1e-6,
365 bias_correction: true,
366 max_compression_error: 1e-4,
367 };
368 Self::with_config(config)
369 }
370
371 pub fn with_config(config: MicroAdamConfig) -> Self {
373 Self {
374 config,
375 state: OptimizerState::new(),
376 momentum: HashMap::new(),
377 variance: HashMap::new(),
378 compression_stats: CompressionStats::default(),
379 }
380 }
381
382 pub fn memory_savings_ratio(&self) -> f32 {
384 if self.compression_stats.total_parameters > 0 {
385 1.0 - (self.compression_stats.total_compressed_size as f32
386 / (self.compression_stats.total_parameters * 2) as f32)
387 } else {
388 0.0
389 }
390 }
391
392 pub fn compression_statistics(&self) -> String {
394 format!(
395 "MicroAdam Compression Stats:\n\
396 - Total parameters: {}\n\
397 - Compressed size: {}\n\
398 - Memory savings: {:.1}%\n\
399 - Average compression ratio: {:.3}\n\
400 - Average compression error: {:.2e}",
401 self.compression_stats.total_parameters,
402 self.compression_stats.total_compressed_size,
403 self.memory_savings_ratio() * 100.0,
404 self.compression_stats.average_compression_ratio,
405 self.compression_stats.average_compression_error
406 )
407 }
408
409 fn update_compression_stats(
411 &mut self,
412 _param_id: &str,
413 compressed: &CompressedGradient,
414 original_gradient: &[f32],
415 ) {
416 self.compression_stats.total_parameters += compressed.original_size;
417 self.compression_stats.total_compressed_size += compressed.compressed_data.len();
418
419 let compression_ratio = compressed.compression_ratio();
420 let compression_error = compressed.compression_error(original_gradient);
421
422 let total_params = self.compression_stats.total_parameters as f32;
424 self.compression_stats.average_compression_ratio =
425 (self.compression_stats.average_compression_ratio
426 * (total_params - compressed.original_size as f32)
427 + compression_ratio * compressed.original_size as f32)
428 / total_params;
429
430 self.compression_stats.average_compression_error =
431 (self.compression_stats.average_compression_error
432 * (total_params - compressed.original_size as f32)
433 + compression_error * compressed.original_size as f32)
434 / total_params;
435
436 let method_name = format!("{:?}", compressed.compression_type);
438 *self.compression_stats.compression_method_usage.entry(method_name).or_insert(0) += 1;
439 }
440}
441
442impl Default for MicroAdam {
443 fn default() -> Self {
444 Self::new()
445 }
446}
447
448impl Optimizer for MicroAdam {
449 fn update(&mut self, parameter: &mut Tensor, grad: &Tensor) -> Result<()> {
450 let param_id = self.state.param_key_for_tensor(parameter)?;
453
454 let grad_data = grad.data()?;
456
457 let compressed_gradient = CompressedGradient::compress(&grad_data, &self.config);
459
460 let compression_error = compressed_gradient.compression_error(&grad_data);
462 if compression_error > self.config.max_compression_error {
463 return Err(TrustformersError::tensor_op_error(
464 &format!(
465 "Compression error {} exceeds maximum allowed {}",
466 compression_error, self.config.max_compression_error
467 ),
468 "MicroAdam::update",
469 ));
470 }
471
472 self.update_compression_stats(¶m_id, &compressed_gradient, &grad_data);
474
475 let momentum = self.momentum.entry(param_id.clone()).or_insert_with(|| {
477 CompressedGradient::compress(&vec![0.0; grad_data.len()], &self.config)
478 });
479
480 let variance = self.variance.entry(param_id.clone()).or_insert_with(|| {
482 CompressedGradient::compress(&vec![0.0; grad_data.len()], &self.config)
483 });
484
485 let mut m = momentum.decompress();
487 let mut v = variance.decompress();
488
489 m.resize(grad_data.len(), 0.0);
491 v.resize(grad_data.len(), 0.0);
492
493 self.state.step();
495
496 let bias_correction1 = if self.config.bias_correction {
498 1.0 - self.config.beta1.powf(self.state.step as f32)
499 } else {
500 1.0
501 };
502
503 let bias_correction2 = if self.config.bias_correction {
504 1.0 - self.config.beta2.powf(self.state.step as f32)
505 } else {
506 1.0
507 };
508
509 for i in 0..grad_data.len() {
511 m[i] = self.config.beta1 * m[i] + (1.0 - self.config.beta1) * grad_data[i];
512 }
513
514 for i in 0..grad_data.len() {
516 v[i] = self.config.beta2 * v[i] + (1.0 - self.config.beta2) * grad_data[i].powi(2);
517 }
518
519 let mut param_data = parameter.data()?;
521 for i in 0..grad_data.len() {
522 let m_hat = m[i] / bias_correction1;
523 let v_hat = v[i] / bias_correction2;
524 let update_val =
525 self.config.learning_rate * m_hat / (v_hat.sqrt() + self.config.epsilon);
526
527 if self.config.weight_decay > 0.0 {
529 param_data[i] *= 1.0 - self.config.learning_rate * self.config.weight_decay;
530 }
531
532 param_data[i] -= update_val;
534 }
535
536 *parameter = Tensor::new(param_data)?;
538
539 *momentum = CompressedGradient::compress(&m, &self.config);
541 *variance = CompressedGradient::compress(&v, &self.config);
542
543 Ok(())
544 }
545
546 fn zero_grad(&mut self) {
547 }
550
551 fn step(&mut self) {
552 }
554
555 fn get_lr(&self) -> f32 {
556 self.config.learning_rate
557 }
558
559 fn set_lr(&mut self, lr: f32) {
560 self.config.learning_rate = lr;
561 }
562}
563
564impl StatefulOptimizer for MicroAdam {
565 type Config = MicroAdamConfig;
566 type State = OptimizerState;
567
568 fn config(&self) -> &Self::Config {
569 &self.config
570 }
571
572 fn state(&self) -> &Self::State {
573 &self.state
574 }
575
576 fn state_mut(&mut self) -> &mut Self::State {
577 &mut self.state
578 }
579
580 fn state_dict(&self) -> Result<HashMap<String, Tensor>> {
581 let mut state_dict = HashMap::new();
582
583 for (param_id, momentum) in &self.momentum {
585 let key = format!("momentum.{}", param_id);
586 let tensor = Tensor::new(momentum.decompress())?;
587 state_dict.insert(key, tensor);
588 }
589
590 for (param_id, variance) in &self.variance {
592 let key = format!("variance.{}", param_id);
593 let tensor = Tensor::new(variance.decompress())?;
594 state_dict.insert(key, tensor);
595 }
596
597 state_dict.insert(
599 "step".to_string(),
600 Tensor::new(vec![self.state.step as f32])?,
601 );
602
603 Ok(state_dict)
604 }
605
606 fn load_state_dict(&mut self, state_dict: HashMap<String, Tensor>) -> Result<()> {
607 if let Some(step_tensor) = state_dict.get("step") {
609 let step_data = step_tensor.data()?;
610 if !step_data.is_empty() {
611 self.state.step = step_data[0] as usize;
612 }
613 }
614
615 for (key, tensor) in &state_dict {
617 if let Some(param_id) = key.strip_prefix("momentum.") {
618 let values = tensor.data()?;
619 let compressed = CompressedGradient::compress(&values, &self.config);
620 self.momentum.insert(param_id.to_string(), compressed);
621 } else if let Some(param_id) = key.strip_prefix("variance.") {
622 let values = tensor.data()?;
623 let compressed = CompressedGradient::compress(&values, &self.config);
624 self.variance.insert(param_id.to_string(), compressed);
625 }
626 }
627
628 Ok(())
629 }
630
631 fn memory_usage(&self) -> StateMemoryStats {
632 let momentum_size: usize = self.momentum.values().map(|m| m.compressed_data.len()).sum();
633 let variance_size: usize = self.variance.values().map(|v| v.compressed_data.len()).sum();
634
635 StateMemoryStats {
636 momentum_elements: momentum_size,
637 variance_elements: variance_size,
638 third_moment_elements: 0,
639 total_bytes: (momentum_size + variance_size) * std::mem::size_of::<f32>(),
640 num_parameters: self.momentum.len(),
641 }
642 }
643
644 fn reset_state(&mut self) {
645 self.state.clear();
646 self.momentum.clear();
647 self.variance.clear();
648 self.compression_stats = CompressionStats::default();
649 }
650
651 fn num_parameters(&self) -> usize {
652 self.momentum.len()
653 }
654}
655
656#[cfg(test)]
657mod tests {
658 use super::*;
659
660 #[test]
661 fn test_microadam_creation() {
662 let optimizer = MicroAdam::new();
663 assert_eq!(optimizer.config.learning_rate, 1e-3);
664 assert_eq!(optimizer.config.beta1, 0.9);
665 assert_eq!(optimizer.config.beta2, 0.999);
666 }
668
669 #[test]
670 fn test_microadam_with_config() {
671 let config = MicroAdamConfig {
672 learning_rate: 2e-3,
673 compression_ratio: 0.2,
674 ..Default::default()
675 };
676 let optimizer = MicroAdam::with_config(config);
677 assert_eq!(optimizer.config.learning_rate, 2e-3);
678 assert_eq!(optimizer.config.compression_ratio, 0.2);
679 }
680
681 #[test]
682 fn test_microadam_for_large_models() {
683 let optimizer = MicroAdam::for_large_models();
684 assert_eq!(optimizer.config.learning_rate, 1e-4);
685 assert_eq!(optimizer.config.compression_ratio, 0.05);
686 assert_eq!(optimizer.config.min_block_size, 128);
687 assert!(optimizer.config.adaptive_compression);
688 }
689
690 #[test]
691 fn test_microadam_for_memory_constrained() {
692 let optimizer = MicroAdam::for_memory_constrained();
693 assert_eq!(optimizer.config.compression_ratio, 0.02);
694 assert_eq!(optimizer.config.min_block_size, 32);
695 assert!(optimizer.config.adaptive_compression);
696 }
697
698 #[test]
699 fn test_compressed_gradient_topk() {
700 let gradient = vec![0.1, 0.05, 0.2, 0.01, 0.15, 0.03];
701 let _config = MicroAdamConfig::default();
702 let compressed = CompressedGradient::compress_topk(&gradient, 3);
703
704 assert_eq!(compressed.compressed_data.len(), 3);
705 assert_eq!(compressed.indices.len(), 3);
706 assert_eq!(compressed.original_size, 6);
707
708 let mut expected_indices = vec![2, 4, 0];
710 let mut actual_indices = compressed.indices.clone();
711 expected_indices.sort();
712 actual_indices.sort();
713 assert_eq!(actual_indices, expected_indices);
714 }
715
716 #[test]
717 fn test_compressed_gradient_threshold() {
718 let gradient = vec![0.1, 0.001, 0.2, 0.0001, 0.15, 0.0003];
719 let config = MicroAdamConfig {
720 compression_threshold: 0.05,
721 ..Default::default()
722 };
723 let compressed = CompressedGradient::compress_threshold(&gradient, &config);
724
725 assert_eq!(compressed.compressed_data.len(), 3);
727 assert_eq!(compressed.indices.len(), 3);
728
729 let mut expected_indices = vec![0, 2, 4];
730 let mut actual_indices = compressed.indices.clone();
731 expected_indices.sort();
732 actual_indices.sort();
733 assert_eq!(actual_indices, expected_indices);
734 }
735
736 #[test]
737 fn test_compression_decompress_cycle() {
738 let gradient = vec![0.1, 0.05, 0.2, 0.01, 0.15, 0.03];
739 let config = MicroAdamConfig::default();
740 let compressed = CompressedGradient::compress(&gradient, &config);
741 let decompressed = compressed.decompress();
742
743 assert_eq!(decompressed.len(), gradient.len());
744
745 for (i, &original) in gradient.iter().enumerate() {
747 if original.abs() > 0.08 {
748 assert!(
750 decompressed[i].abs() > 0.0,
751 "Significant value at index {} was lost",
752 i
753 );
754 }
755 }
756 }
757
758 #[test]
759 fn test_compression_error_calculation() {
760 let gradient = vec![0.1, 0.05, 0.2, 0.01, 0.15, 0.03];
761 let config = MicroAdamConfig::default();
762 let compressed = CompressedGradient::compress(&gradient, &config);
763 let error = compressed.compression_error(&gradient);
764
765 assert!(error >= 0.0);
766 assert!(error <= 1.0); }
768
769 #[test]
770 fn test_microadam_update() -> Result<()> {
771 let mut optimizer = MicroAdam::new();
772 let gradient_data = vec![0.1, -0.05, 0.2, -0.01];
773 let gradient = Tensor::new(gradient_data.clone())?;
774 let mut parameter = Tensor::new(vec![1.0, 1.0, 1.0, 1.0])?;
775
776 optimizer.update(&mut parameter, &gradient)?;
777
778 assert_eq!(optimizer.state().step, 1);
780
781 let param_data = parameter.data()?;
783 assert_eq!(param_data.len(), gradient_data.len());
784
785 assert_ne!(param_data[0], 1.0);
787
788 Ok(())
789 }
790
791 #[test]
792 fn test_microadam_multiple_updates() -> Result<()> {
793 let mut optimizer = MicroAdam::new();
794 let gradient_data = vec![0.1, -0.05, 0.2, -0.01];
795 let gradient = Tensor::new(gradient_data)?;
796 let mut parameter = Tensor::new(vec![1.0, 1.0, 1.0, 1.0])?;
797
798 for i in 1..=5 {
800 optimizer.update(&mut parameter, &gradient)?;
801 assert_eq!(optimizer.state().step, i);
802 }
803
804 Ok(())
805 }
806
807 #[test]
808 fn test_memory_savings_ratio() {
809 let config = MicroAdamConfig {
810 max_compression_error: 1.0, ..MicroAdamConfig::default()
812 };
813 let mut optimizer = MicroAdam::with_config(config);
814
815 assert_eq!(optimizer.memory_savings_ratio(), 0.0);
817
818 let gradient_data = vec![0.1; 1000]; let gradient = Tensor::new(gradient_data).expect("Failed to create tensor");
821 let mut parameter = Tensor::new(vec![1.0; 1000]).expect("Failed to create tensor");
822 optimizer.update(&mut parameter, &gradient).expect("Optimizer update failed");
823
824 let savings = optimizer.memory_savings_ratio();
825 assert!(savings > 0.0, "Should show memory savings");
826 assert!(savings < 1.0, "Savings ratio should be less than 100%");
827 }
828
829 #[test]
830 fn test_compression_statistics() {
831 let config = MicroAdamConfig {
832 max_compression_error: 1.0, ..MicroAdamConfig::default()
834 };
835 let mut optimizer = MicroAdam::with_config(config);
836 let gradient_data = vec![0.1; 500];
837 let gradient = Tensor::new(gradient_data).expect("Failed to create tensor");
838 let mut parameter = Tensor::new(vec![1.0; 500]).expect("Failed to create tensor");
839
840 optimizer.update(&mut parameter, &gradient).expect("Optimizer update failed");
841
842 let stats = optimizer.compression_statistics();
843 assert!(stats.contains("MicroAdam Compression Stats"));
844 assert!(stats.contains("Total parameters: 500"));
845 assert!(stats.contains("Memory savings"));
846 assert!(stats.contains("compression ratio"));
847 }
848
849 #[test]
850 fn test_learning_rate_setter_getter() {
851 let mut optimizer = MicroAdam::new();
852 assert_eq!(optimizer.get_lr(), 1e-3);
853
854 optimizer.set_lr(2e-3);
855 assert_eq!(optimizer.get_lr(), 2e-3);
856 }
857
858 #[test]
859 fn test_state_dict_operations() -> Result<()> {
860 let mut optimizer = MicroAdam::new();
861 let gradient_data = vec![0.1, -0.05, 0.2];
862 let gradient = Tensor::new(gradient_data)?;
863 let mut param1 = Tensor::new(vec![1.0, 1.0, 1.0])?;
864 let mut param2 = Tensor::new(vec![2.0, 2.0, 2.0])?;
865
866 optimizer.update(&mut param1, &gradient)?;
868 optimizer.update(&mut param2, &gradient)?;
869
870 let state_dict = optimizer.state_dict()?;
872 assert!(state_dict.contains_key("step"));
873
874 let mut new_optimizer = MicroAdam::new();
876 new_optimizer.load_state_dict(state_dict)?;
877
878 assert_eq!(new_optimizer.state().step, optimizer.state().step);
879
880 Ok(())
881 }
882
883 #[test]
884 fn test_memory_usage_tracking() -> Result<()> {
885 let config = MicroAdamConfig {
886 max_compression_error: 1.0, ..MicroAdamConfig::default()
888 };
889 let mut optimizer = MicroAdam::with_config(config);
890 let initial_usage = optimizer.memory_usage();
891
892 let gradient_data = vec![0.1; 1000];
893 let gradient = Tensor::new(gradient_data)?;
894 let mut parameter = Tensor::new(vec![1.0; 1000])?;
895 optimizer.update(&mut parameter, &gradient)?;
896
897 let after_usage = optimizer.memory_usage();
898 assert!(after_usage.total_bytes > initial_usage.total_bytes);
899 assert!(after_usage.momentum_elements > 0);
900 assert!(after_usage.variance_elements > 0);
901
902 Ok(())
903 }
904
905 #[test]
906 fn test_adaptive_compression_selection() {
907 let sparse_gradient = vec![0.0; 1000]; let dense_gradient = vec![0.1; 1000]; let config = MicroAdamConfig {
911 adaptive_compression: true,
912 compression_threshold: 1e-6,
913 ..Default::default()
914 };
915
916 let sparse_compression =
917 CompressedGradient::choose_adaptive_compression(&sparse_gradient, &config);
918 let dense_compression =
919 CompressedGradient::choose_adaptive_compression(&dense_gradient, &config);
920
921 match sparse_compression {
924 CompressionType::Threshold
925 | CompressionType::TopK
926 | CompressionType::BlockWise
927 | CompressionType::Adaptive => {},
928 }
929
930 match dense_compression {
931 CompressionType::Threshold
932 | CompressionType::TopK
933 | CompressionType::BlockWise
934 | CompressionType::Adaptive => {},
935 }
936 }
937}