1use parking_lot::Mutex;
7use scirs2_core::numeric::{FromPrimitive, ToPrimitive};
8use std::collections::HashMap;
9use std::sync::{Arc, RwLock};
10use torsh_core::dtype::FloatElement;
11use torsh_core::error::{Result, TorshError};
12use torsh_core::sync::RwLockExt;
13
14#[derive(Debug, Clone)]
16pub struct CompressionConfig {
17 pub algorithm: CompressionAlgorithm,
19 pub target_ratio: f64,
21 pub error_tolerance: f64,
23 pub memory_budget: usize,
25 pub adaptive: bool,
27 pub sparsity_threshold: f64,
29}
30
31impl Default for CompressionConfig {
32 fn default() -> Self {
33 Self {
34 algorithm: CompressionAlgorithm::Quantization8Bit,
35 target_ratio: 0.25, error_tolerance: 1e-4,
37 memory_budget: 100 * 1024 * 1024, adaptive: true,
39 sparsity_threshold: 0.01, }
41 }
42}
43
44#[derive(Debug, Clone, Copy, PartialEq, Eq)]
46pub enum CompressionAlgorithm {
47 None,
49 Quantization8Bit,
51 Quantization4Bit,
53 Quantization2Bit,
55 Quantization1Bit,
57 TopKSparsification,
59 RandomSparsification,
61 GradientSketching,
63 PowerSGD,
65 ErrorFeedback,
67 Adaptive,
69}
70
71#[derive(Debug, Clone)]
73pub struct CompressedGradient {
74 pub original_shape: Vec<usize>,
76 pub data: Vec<u8>,
78 pub metadata: CompressionMetadata,
80 pub algorithm: CompressionAlgorithm,
82}
83
84#[derive(Debug, Clone)]
86pub struct CompressionMetadata {
87 pub scale: f64,
89 pub zero_point: i32,
91 pub indices: Vec<usize>,
93 pub seed: u64,
95 pub rank: usize,
97 pub error: Vec<f64>,
99}
100
101impl Default for CompressionMetadata {
102 fn default() -> Self {
103 Self {
104 scale: 1.0,
105 zero_point: 0,
106 indices: Vec::new(),
107 seed: 0,
108 rank: 0,
109 error: Vec::new(),
110 }
111 }
112}
113
114#[derive(Clone)]
116pub struct GradientCompressor<T: FloatElement> {
117 config: CompressionConfig,
119 error_feedback: Arc<RwLock<HashMap<String, Vec<T>>>>,
121 stats: Arc<RwLock<CompressionStats>>,
123 rng_state: Arc<Mutex<u64>>,
125}
126
127impl<T: FloatElement> std::fmt::Debug for GradientCompressor<T> {
128 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
129 f.debug_struct("GradientCompressor")
130 .field("config", &self.config)
131 .field(
132 "error_feedback_size",
133 &self.error_feedback.read_or_recover().len(),
134 )
135 .field("stats", &self.stats.read_or_recover())
136 .finish()
137 }
138}
139
140#[derive(Debug, Clone, Default)]
142pub struct CompressionStats {
143 pub total_compressions: usize,
145 pub total_bytes_original: usize,
147 pub total_bytes_compressed: usize,
149 pub avg_compression_ratio: f64,
151 pub total_compression_time_ms: u64,
153 pub total_decompression_time_ms: u64,
155 pub compression_error: f64,
157}
158
159impl<T: FloatElement + FromPrimitive + ToPrimitive> GradientCompressor<T> {
160 pub fn new(config: CompressionConfig) -> Self {
162 Self {
163 config,
164 error_feedback: Arc::new(RwLock::new(HashMap::new())),
165 stats: Arc::new(RwLock::new(CompressionStats::default())),
166 rng_state: Arc::new(Mutex::new(42)), }
168 }
169
170 pub fn compress(
172 &mut self,
173 gradients: &[T],
174 parameter_name: &str,
175 ) -> Result<CompressedGradient> {
176 let start_time = std::time::Instant::now();
177
178 let algorithm = if self.config.adaptive {
179 self.choose_best_algorithm(gradients)?
180 } else {
181 self.config.algorithm
182 };
183
184 let compressed = match algorithm {
185 CompressionAlgorithm::None => self.compress_none(gradients)?,
186 CompressionAlgorithm::Quantization8Bit => self.compress_quantization_8bit(gradients)?,
187 CompressionAlgorithm::Quantization4Bit => self.compress_quantization_4bit(gradients)?,
188 CompressionAlgorithm::Quantization2Bit => self.compress_quantization_2bit(gradients)?,
189 CompressionAlgorithm::Quantization1Bit => self.compress_quantization_1bit(gradients)?,
190 CompressionAlgorithm::TopKSparsification => {
191 self.compress_top_k_sparsification(gradients)?
192 }
193 CompressionAlgorithm::RandomSparsification => {
194 self.compress_random_sparsification(gradients)?
195 }
196 CompressionAlgorithm::GradientSketching => {
197 self.compress_gradient_sketching(gradients)?
198 }
199 CompressionAlgorithm::PowerSGD => self.compress_power_sgd(gradients)?,
200 CompressionAlgorithm::ErrorFeedback => {
201 self.compress_error_feedback(gradients, parameter_name)?
202 }
203 CompressionAlgorithm::Adaptive => {
204 return Err(TorshError::AutogradError(
205 "Adaptive algorithm should have been resolved".to_string(),
206 ));
207 }
208 };
209
210 let compression_time = start_time.elapsed().as_millis() as u64;
212 let mut stats = self.stats.write_or_recover();
213 stats.total_compressions += 1;
214 stats.total_bytes_original += std::mem::size_of_val(gradients);
215 stats.total_bytes_compressed += compressed.data.len();
216 stats.total_compression_time_ms += compression_time;
217
218 let compression_ratio =
219 compressed.data.len() as f64 / std::mem::size_of_val(gradients) as f64;
220 stats.avg_compression_ratio = (stats.avg_compression_ratio
221 * (stats.total_compressions - 1) as f64
222 + compression_ratio)
223 / stats.total_compressions as f64;
224
225 Ok(compressed)
226 }
227
228 pub fn decompress(&mut self, compressed: &CompressedGradient) -> Result<Vec<T>> {
230 let start_time = std::time::Instant::now();
231
232 let decompressed = match compressed.algorithm {
233 CompressionAlgorithm::None => self.decompress_none(compressed)?,
234 CompressionAlgorithm::Quantization8Bit => {
235 self.decompress_quantization_8bit(compressed)?
236 }
237 CompressionAlgorithm::Quantization4Bit => {
238 self.decompress_quantization_4bit(compressed)?
239 }
240 CompressionAlgorithm::Quantization2Bit => {
241 self.decompress_quantization_2bit(compressed)?
242 }
243 CompressionAlgorithm::Quantization1Bit => {
244 self.decompress_quantization_1bit(compressed)?
245 }
246 CompressionAlgorithm::TopKSparsification => {
247 self.decompress_top_k_sparsification(compressed)?
248 }
249 CompressionAlgorithm::RandomSparsification => {
250 self.decompress_random_sparsification(compressed)?
251 }
252 CompressionAlgorithm::GradientSketching => {
253 self.decompress_gradient_sketching(compressed)?
254 }
255 CompressionAlgorithm::PowerSGD => self.decompress_power_sgd(compressed)?,
256 CompressionAlgorithm::ErrorFeedback => self.decompress_error_feedback(compressed)?,
257 CompressionAlgorithm::Adaptive => {
258 return Err(TorshError::AutogradError(
259 "Cannot decompress adaptive algorithm directly".to_string(),
260 ));
261 }
262 };
263
264 let decompression_time = start_time.elapsed().as_millis() as u64;
266 self.stats.write_or_recover().total_decompression_time_ms += decompression_time;
267
268 Ok(decompressed)
269 }
270
271 fn choose_best_algorithm(&self, gradients: &[T]) -> Result<CompressionAlgorithm> {
273 let sparsity = self.calculate_sparsity(gradients);
275 let variance = self.calculate_variance(gradients);
276 let magnitude = self.calculate_magnitude(gradients);
277
278 if sparsity > self.config.sparsity_threshold {
280 Ok(CompressionAlgorithm::TopKSparsification)
281 } else if variance < 0.01 && magnitude < 1.0 {
282 Ok(CompressionAlgorithm::Quantization2Bit)
283 } else if variance < 0.1 {
284 Ok(CompressionAlgorithm::Quantization4Bit)
285 } else if gradients.len() > 10000 {
286 Ok(CompressionAlgorithm::PowerSGD)
287 } else {
288 Ok(CompressionAlgorithm::Quantization8Bit)
289 }
290 }
291
292 fn calculate_sparsity(&self, gradients: &[T]) -> f64 {
294 let threshold = <T as torsh_core::dtype::TensorElement>::from_f64(1e-8)
295 .expect("f64 conversion should succeed");
296 let near_zero_count = gradients.iter().filter(|&&x| x.abs() < threshold).count();
297 near_zero_count as f64 / gradients.len() as f64
298 }
299
300 fn calculate_variance(&self, gradients: &[T]) -> f64 {
302 if gradients.is_empty() {
303 return 0.0;
304 }
305
306 let mean = gradients
307 .iter()
308 .map(|x| ToPrimitive::to_f64(x).expect("f64 conversion should succeed"))
309 .sum::<f64>()
310 / gradients.len() as f64;
311
312 let variance = gradients
313 .iter()
314 .map(|x| {
315 let val = ToPrimitive::to_f64(x).expect("f64 conversion should succeed");
316 (val - mean).powi(2)
317 })
318 .sum::<f64>()
319 / gradients.len() as f64;
320
321 variance
322 }
323
324 fn calculate_magnitude(&self, gradients: &[T]) -> f64 {
326 if gradients.is_empty() {
327 return 0.0;
328 }
329
330 let sum_squares = gradients
331 .iter()
332 .map(|x| {
333 ToPrimitive::to_f64(x)
334 .expect("f64 conversion should succeed")
335 .powi(2)
336 })
337 .sum::<f64>();
338
339 (sum_squares / gradients.len() as f64).sqrt()
340 }
341
342 fn compress_none(&self, gradients: &[T]) -> Result<CompressedGradient> {
344 let data = unsafe {
345 std::slice::from_raw_parts(
346 gradients.as_ptr() as *const u8,
347 std::mem::size_of_val(gradients),
348 )
349 .to_vec()
350 };
351
352 Ok(CompressedGradient {
353 original_shape: vec![gradients.len()],
354 data,
355 metadata: CompressionMetadata::default(),
356 algorithm: CompressionAlgorithm::None,
357 })
358 }
359
360 fn compress_quantization_8bit(&self, gradients: &[T]) -> Result<CompressedGradient> {
362 let mut min_val = f64::INFINITY;
364 let mut max_val = f64::NEG_INFINITY;
365
366 for &grad in gradients {
367 let val = ToPrimitive::to_f64(&grad).expect("f64 conversion should succeed");
368 min_val = min_val.min(val);
369 max_val = max_val.max(val);
370 }
371
372 let scale = (max_val - min_val) / 255.0;
374 let zero_point = (-min_val / scale).round() as i32;
375
376 let mut quantized = Vec::with_capacity(gradients.len());
378 for &grad in gradients {
379 let val = ToPrimitive::to_f64(&grad).expect("f64 conversion should succeed");
380 let quantized_val = ((val / scale) + zero_point as f64).round() as u8;
381 quantized.push(quantized_val);
382 }
383
384 let metadata = CompressionMetadata {
385 scale,
386 zero_point,
387 ..Default::default()
388 };
389
390 Ok(CompressedGradient {
391 original_shape: vec![gradients.len()],
392 data: quantized,
393 metadata,
394 algorithm: CompressionAlgorithm::Quantization8Bit,
395 })
396 }
397
398 fn compress_quantization_4bit(&self, gradients: &[T]) -> Result<CompressedGradient> {
400 let mut min_val = f64::INFINITY;
402 let mut max_val = f64::NEG_INFINITY;
403
404 for &grad in gradients {
405 let val = ToPrimitive::to_f64(&grad).expect("f64 conversion should succeed");
406 min_val = min_val.min(val);
407 max_val = max_val.max(val);
408 }
409
410 let scale = (max_val - min_val) / 15.0; let zero_point = (-min_val / scale).round() as i32;
412
413 let mut quantized = Vec::with_capacity(gradients.len().div_ceil(2));
415 for chunk in gradients.chunks(2) {
416 let first = if !chunk.is_empty() {
417 let val = ToPrimitive::to_f64(&chunk[0]).expect("f64 conversion should succeed");
418 ((val / scale) + zero_point as f64).round().clamp(0.0, 15.0) as u8
419 } else {
420 0
421 };
422
423 let second = if chunk.len() > 1 {
424 let val = ToPrimitive::to_f64(&chunk[1]).expect("f64 conversion should succeed");
425 ((val / scale) + zero_point as f64).round().clamp(0.0, 15.0) as u8
426 } else {
427 0
428 };
429
430 quantized.push((first << 4) | second);
431 }
432
433 let metadata = CompressionMetadata {
434 scale,
435 zero_point,
436 ..Default::default()
437 };
438
439 Ok(CompressedGradient {
440 original_shape: vec![gradients.len()],
441 data: quantized,
442 metadata,
443 algorithm: CompressionAlgorithm::Quantization4Bit,
444 })
445 }
446
447 fn compress_quantization_2bit(&self, gradients: &[T]) -> Result<CompressedGradient> {
449 let mut min_val = f64::INFINITY;
450 let mut max_val = f64::NEG_INFINITY;
451
452 for &grad in gradients {
453 let val = ToPrimitive::to_f64(&grad).expect("f64 conversion should succeed");
454 min_val = min_val.min(val);
455 max_val = max_val.max(val);
456 }
457
458 let scale = (max_val - min_val) / 3.0; let zero_point = (-min_val / scale).round() as i32;
460
461 let mut quantized = Vec::with_capacity(gradients.len().div_ceil(4));
463 for chunk in gradients.chunks(4) {
464 let mut byte_val = 0u8;
465 for (i, &val) in chunk.iter().enumerate() {
466 let val_f64 = ToPrimitive::to_f64(&val).expect("f64 conversion should succeed");
467 let quantized_val = ((val_f64 / scale) + zero_point as f64)
468 .round()
469 .clamp(0.0, 3.0) as u8;
470 byte_val |= quantized_val << (i * 2);
471 }
472 quantized.push(byte_val);
473 }
474
475 let metadata = CompressionMetadata {
476 scale,
477 zero_point,
478 ..Default::default()
479 };
480
481 Ok(CompressedGradient {
482 original_shape: vec![gradients.len()],
483 data: quantized,
484 metadata,
485 algorithm: CompressionAlgorithm::Quantization2Bit,
486 })
487 }
488
489 fn compress_quantization_1bit(&self, gradients: &[T]) -> Result<CompressedGradient> {
491 let magnitude = self.calculate_magnitude(gradients);
493
494 let mut quantized = Vec::with_capacity(gradients.len().div_ceil(8));
496 for chunk in gradients.chunks(8) {
497 let mut byte_val = 0u8;
498 for (i, &val) in chunk.iter().enumerate() {
499 let val_f64 = ToPrimitive::to_f64(&val).expect("f64 conversion should succeed");
500 if val_f64 >= 0.0 {
501 byte_val |= 1 << i;
502 }
503 }
504 quantized.push(byte_val);
505 }
506
507 let metadata = CompressionMetadata {
508 scale: magnitude,
509 ..Default::default()
510 };
511
512 Ok(CompressedGradient {
513 original_shape: vec![gradients.len()],
514 data: quantized,
515 metadata,
516 algorithm: CompressionAlgorithm::Quantization1Bit,
517 })
518 }
519
520 fn compress_top_k_sparsification(&self, gradients: &[T]) -> Result<CompressedGradient> {
522 let k = (gradients.len() as f64 * (1.0 - self.config.sparsity_threshold)).round() as usize;
523
524 let mut value_index_pairs: Vec<(f64, usize)> = gradients
526 .iter()
527 .enumerate()
528 .map(|(i, &val)| {
529 (
530 ToPrimitive::to_f64(&val)
531 .expect("f64 conversion should succeed")
532 .abs(),
533 i,
534 )
535 })
536 .collect();
537
538 value_index_pairs
539 .sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
540
541 let top_k_indices: Vec<usize> = value_index_pairs
543 .iter()
544 .take(k)
545 .map(|(_, idx)| *idx)
546 .collect();
547
548 let mut compressed_data = Vec::new();
550 let mut indices = Vec::new();
551
552 for &idx in &top_k_indices {
553 let val = ToPrimitive::to_f64(&gradients[idx]).expect("f64 conversion should succeed");
554 let bytes = val.to_le_bytes();
555 compressed_data.extend_from_slice(&bytes);
556 indices.push(idx);
557 }
558
559 let metadata = CompressionMetadata {
560 indices,
561 ..Default::default()
562 };
563
564 Ok(CompressedGradient {
565 original_shape: vec![gradients.len()],
566 data: compressed_data,
567 metadata,
568 algorithm: CompressionAlgorithm::TopKSparsification,
569 })
570 }
571
572 fn compress_random_sparsification(&self, gradients: &[T]) -> Result<CompressedGradient> {
574 let keep_prob = 1.0 - self.config.sparsity_threshold;
575 let mut rng_state = self.rng_state.lock();
576
577 let mut compressed_data = Vec::new();
578 let mut indices = Vec::new();
579
580 for (i, &val) in gradients.iter().enumerate() {
581 *rng_state = (1103515245_u64.wrapping_mul(*rng_state).wrapping_add(12345)) % (1 << 31);
583 let random_val = *rng_state as f64 / (1u64 << 31) as f64;
584
585 if random_val < keep_prob {
586 let val_f64 = ToPrimitive::to_f64(&val).expect("f64 conversion should succeed");
587 let bytes = val_f64.to_le_bytes();
588 compressed_data.extend_from_slice(&bytes);
589 indices.push(i);
590 }
591 }
592
593 let metadata = CompressionMetadata {
594 indices,
595 seed: *rng_state,
596 ..Default::default()
597 };
598
599 Ok(CompressedGradient {
600 original_shape: vec![gradients.len()],
601 data: compressed_data,
602 metadata,
603 algorithm: CompressionAlgorithm::RandomSparsification,
604 })
605 }
606
607 fn compress_gradient_sketching(&self, gradients: &[T]) -> Result<CompressedGradient> {
609 let sketch_size = (gradients.len() as f64 * self.config.target_ratio).round() as usize;
611 let sketch_size = sketch_size.max(1);
612
613 let mut rng_state = self.rng_state.lock();
614 let mut sketch = vec![0.0; sketch_size];
615
616 for &grad in gradients {
617 let val = ToPrimitive::to_f64(&grad).expect("f64 conversion should succeed");
618
619 for j in 0..sketch_size {
621 *rng_state =
622 (1103515245_u64.wrapping_mul(*rng_state).wrapping_add(12345)) % (1 << 31);
623 let random_sign = if (*rng_state % 2) == 0 { 1.0 } else { -1.0 };
624 sketch[j] += val * random_sign;
625 }
626 }
627
628 let mut compressed_data = Vec::new();
630 for &val in &sketch {
631 let bytes = val.to_le_bytes();
632 compressed_data.extend_from_slice(&bytes);
633 }
634
635 let metadata = CompressionMetadata {
636 seed: *rng_state,
637 rank: sketch_size,
638 ..Default::default()
639 };
640
641 Ok(CompressedGradient {
642 original_shape: vec![gradients.len()],
643 data: compressed_data,
644 metadata,
645 algorithm: CompressionAlgorithm::GradientSketching,
646 })
647 }
648
649 fn compress_power_sgd(&self, gradients: &[T]) -> Result<CompressedGradient> {
651 let rank = (gradients.len() as f64 * self.config.target_ratio)
653 .sqrt()
654 .round() as usize;
655 let rank = rank.max(1).min(gradients.len());
656
657 let mut compressed_data = Vec::new();
660
661 for i in 0..rank.min(gradients.len()) {
663 let val = ToPrimitive::to_f64(&gradients[i]).expect("f64 conversion should succeed");
664 let bytes = val.to_le_bytes();
665 compressed_data.extend_from_slice(&bytes);
666 }
667
668 let metadata = CompressionMetadata {
669 rank,
670 ..Default::default()
671 };
672
673 Ok(CompressedGradient {
674 original_shape: vec![gradients.len()],
675 data: compressed_data,
676 metadata,
677 algorithm: CompressionAlgorithm::PowerSGD,
678 })
679 }
680
681 fn compress_error_feedback(
683 &mut self,
684 gradients: &[T],
685 parameter_name: &str,
686 ) -> Result<CompressedGradient> {
687 let mut error_feedback = self.error_feedback.write_or_recover();
689 let error_buffer = error_feedback
690 .entry(parameter_name.to_string())
691 .or_insert_with(|| {
692 vec![<T as torsh_core::dtype::TensorElement>::zero(); gradients.len()]
693 });
694
695 if error_buffer.len() != gradients.len() {
697 error_buffer.resize(
698 gradients.len(),
699 <T as torsh_core::dtype::TensorElement>::zero(),
700 );
701 }
702
703 let mut compensated_gradients = Vec::with_capacity(gradients.len());
705 for (&grad, &error) in gradients.iter().zip(error_buffer.iter()) {
706 compensated_gradients.push(grad + error);
707 }
708
709 let compressed = self.compress_quantization_8bit(&compensated_gradients)?;
711
712 let decompressed = self.decompress_quantization_8bit(&compressed)?;
714 for (i, (&original, &decompressed_val)) in compensated_gradients
715 .iter()
716 .zip(decompressed.iter())
717 .enumerate()
718 {
719 let error = original - decompressed_val;
720 error_buffer[i] = error;
721 }
722
723 Ok(CompressedGradient {
724 algorithm: CompressionAlgorithm::ErrorFeedback,
725 ..compressed
726 })
727 }
728
729 fn decompress_none(&self, compressed: &CompressedGradient) -> Result<Vec<T>> {
731 let gradients = unsafe {
732 std::slice::from_raw_parts(
733 compressed.data.as_ptr() as *const T,
734 compressed.data.len() / std::mem::size_of::<T>(),
735 )
736 .to_vec()
737 };
738 Ok(gradients)
739 }
740
741 fn decompress_quantization_8bit(&self, compressed: &CompressedGradient) -> Result<Vec<T>> {
742 let scale = compressed.metadata.scale;
743 let zero_point = compressed.metadata.zero_point;
744
745 let mut gradients = Vec::with_capacity(compressed.data.len());
746 for &quantized_val in &compressed.data {
747 let dequantized = (quantized_val as f64 - zero_point as f64) * scale;
748 gradients.push(
749 <T as torsh_core::dtype::TensorElement>::from_f64(dequantized)
750 .expect("f64 conversion should succeed"),
751 );
752 }
753
754 Ok(gradients)
755 }
756
757 fn decompress_quantization_4bit(&self, compressed: &CompressedGradient) -> Result<Vec<T>> {
758 let scale = compressed.metadata.scale;
759 let zero_point = compressed.metadata.zero_point;
760 let original_size = compressed.original_shape[0];
761
762 let mut gradients = Vec::with_capacity(original_size);
763 for &byte_val in &compressed.data {
764 let first = (byte_val >> 4) & 0x0F;
766 let dequantized_first = (first as f64 - zero_point as f64) * scale;
767 gradients.push(
768 <T as torsh_core::dtype::TensorElement>::from_f64(dequantized_first)
769 .expect("f64 conversion should succeed"),
770 );
771
772 if gradients.len() < original_size {
773 let second = byte_val & 0x0F;
775 let dequantized_second = (second as f64 - zero_point as f64) * scale;
776 gradients.push(
777 <T as torsh_core::dtype::TensorElement>::from_f64(dequantized_second)
778 .expect("f64 conversion should succeed"),
779 );
780 }
781 }
782
783 gradients.truncate(original_size);
784 Ok(gradients)
785 }
786
787 fn decompress_quantization_2bit(&self, compressed: &CompressedGradient) -> Result<Vec<T>> {
788 let scale = compressed.metadata.scale;
789 let zero_point = compressed.metadata.zero_point;
790 let original_size = compressed.original_shape[0];
791
792 let mut gradients = Vec::with_capacity(original_size);
793 for &byte_val in &compressed.data {
794 for i in 0..4 {
795 if gradients.len() >= original_size {
796 break;
797 }
798 let quantized_val = (byte_val >> (i * 2)) & 0x03;
799 let dequantized = (quantized_val as f64 - zero_point as f64) * scale;
800 gradients.push(
801 <T as torsh_core::dtype::TensorElement>::from_f64(dequantized)
802 .expect("f64 conversion should succeed"),
803 );
804 }
805 }
806
807 gradients.truncate(original_size);
808 Ok(gradients)
809 }
810
811 fn decompress_quantization_1bit(&self, compressed: &CompressedGradient) -> Result<Vec<T>> {
812 let magnitude = compressed.metadata.scale;
813 let original_size = compressed.original_shape[0];
814
815 let mut gradients = Vec::with_capacity(original_size);
816 for &byte_val in &compressed.data {
817 for i in 0..8 {
818 if gradients.len() >= original_size {
819 break;
820 }
821 let sign_bit = (byte_val >> i) & 1;
822 let value = if sign_bit == 1 { magnitude } else { -magnitude };
823 gradients.push(
824 <T as torsh_core::dtype::TensorElement>::from_f64(value)
825 .expect("f64 conversion should succeed"),
826 );
827 }
828 }
829
830 gradients.truncate(original_size);
831 Ok(gradients)
832 }
833
834 fn decompress_top_k_sparsification(&self, compressed: &CompressedGradient) -> Result<Vec<T>> {
835 let original_size = compressed.original_shape[0];
836 let mut gradients = vec![<T as torsh_core::dtype::TensorElement>::zero(); original_size];
837
838 let values_per_element = std::mem::size_of::<f64>();
839 let num_values = compressed.data.len() / values_per_element;
840
841 for (i, &idx) in compressed
842 .metadata
843 .indices
844 .iter()
845 .take(num_values)
846 .enumerate()
847 {
848 let start = i * values_per_element;
849 let end = start + values_per_element;
850 if end <= compressed.data.len() && idx < original_size {
851 let bytes = &compressed.data[start..end];
852 let value =
853 f64::from_le_bytes(bytes.try_into().expect("slice should be 8 bytes for f64"));
854 gradients[idx] = <T as torsh_core::dtype::TensorElement>::from_f64(value)
855 .expect("f64 conversion should succeed");
856 }
857 }
858
859 Ok(gradients)
860 }
861
862 fn decompress_random_sparsification(&self, compressed: &CompressedGradient) -> Result<Vec<T>> {
863 let original_size = compressed.original_shape[0];
864 let mut gradients = vec![<T as torsh_core::dtype::TensorElement>::zero(); original_size];
865
866 let values_per_element = std::mem::size_of::<f64>();
867 let num_values = compressed.data.len() / values_per_element;
868
869 for (i, &idx) in compressed
870 .metadata
871 .indices
872 .iter()
873 .take(num_values)
874 .enumerate()
875 {
876 let start = i * values_per_element;
877 let end = start + values_per_element;
878 if end <= compressed.data.len() && idx < original_size {
879 let bytes = &compressed.data[start..end];
880 let value =
881 f64::from_le_bytes(bytes.try_into().expect("slice should be 8 bytes for f64"));
882 gradients[idx] = <T as torsh_core::dtype::TensorElement>::from_f64(value)
883 .expect("f64 conversion should succeed");
884 }
885 }
886
887 Ok(gradients)
888 }
889
890 fn decompress_gradient_sketching(&self, _compressed: &CompressedGradient) -> Result<Vec<T>> {
891 Err(TorshError::AutogradError(
894 "Gradient sketching decompression not implemented - lossy compression".to_string(),
895 ))
896 }
897
898 fn decompress_power_sgd(&self, compressed: &CompressedGradient) -> Result<Vec<T>> {
899 let original_size = compressed.original_shape[0];
900 let rank = compressed.metadata.rank;
901
902 let mut gradients = vec![<T as torsh_core::dtype::TensorElement>::zero(); original_size];
903
904 let values_per_element = std::mem::size_of::<f64>();
906 let num_values = (compressed.data.len() / values_per_element).min(rank);
907
908 for i in 0..num_values.min(original_size) {
909 let start = i * values_per_element;
910 let end = start + values_per_element;
911 let bytes = &compressed.data[start..end];
912 let value =
913 f64::from_le_bytes(bytes.try_into().expect("slice should be 8 bytes for f64"));
914 gradients[i] = <T as torsh_core::dtype::TensorElement>::from_f64(value)
915 .expect("f64 conversion should succeed");
916 }
917
918 Ok(gradients)
919 }
920
921 fn decompress_error_feedback(&self, compressed: &CompressedGradient) -> Result<Vec<T>> {
922 self.decompress_quantization_8bit(compressed)
924 }
925
926 pub fn get_stats(&self) -> CompressionStats {
928 self.stats.read_or_recover().clone()
929 }
930
931 pub fn reset_stats(&mut self) {
933 *self.stats.write_or_recover() = CompressionStats::default();
934 }
935
936 pub fn update_config(&mut self, new_config: CompressionConfig) {
938 self.config = new_config;
939 }
940}
941
942pub mod utils {
944 use super::*;
945
946 pub fn analyze_gradients<T: FloatElement + ToPrimitive>(gradients: &[T]) -> GradientAnalysis {
948 if gradients.is_empty() {
949 return GradientAnalysis::default();
950 }
951
952 let mut min_val = f64::INFINITY;
953 let mut max_val = f64::NEG_INFINITY;
954 let mut sum = 0.0;
955 let mut sum_squares = 0.0;
956 let mut zero_count = 0;
957
958 for &val in gradients {
959 let val_f64 = ToPrimitive::to_f64(&val).expect("f64 conversion should succeed");
960 min_val = min_val.min(val_f64);
961 max_val = max_val.max(val_f64);
962 sum += val_f64;
963 sum_squares += val_f64 * val_f64;
964
965 if val_f64.abs() < 1e-8 {
966 zero_count += 1;
967 }
968 }
969
970 let n = gradients.len() as f64;
971 let mean = sum / n;
972 let variance = (sum_squares / n) - (mean * mean);
973 let std_dev = variance.sqrt();
974 let sparsity = zero_count as f64 / n;
975
976 GradientAnalysis {
977 min_value: min_val,
978 max_value: max_val,
979 mean,
980 std_dev,
981 sparsity,
982 dynamic_range: max_val - min_val,
983 recommended_algorithm: if sparsity > 0.1 {
984 CompressionAlgorithm::TopKSparsification
985 } else if std_dev < 0.01 {
986 CompressionAlgorithm::Quantization2Bit
987 } else if std_dev < 0.1 {
988 CompressionAlgorithm::Quantization4Bit
989 } else {
990 CompressionAlgorithm::Quantization8Bit
991 },
992 }
993 }
994
995 pub fn benchmark_compression<T: FloatElement + FromPrimitive + ToPrimitive>(
997 gradients: &[T],
998 algorithms: &[CompressionAlgorithm],
999 ) -> Vec<CompressionBenchmark> {
1000 let mut results = Vec::new();
1001
1002 for &algorithm in algorithms {
1003 let config = CompressionConfig {
1004 algorithm,
1005 ..Default::default()
1006 };
1007
1008 let mut compressor = GradientCompressor::new(config);
1009
1010 let start_time = std::time::Instant::now();
1011 if let Ok(compressed) = compressor.compress(gradients, "benchmark") {
1012 let compression_time = start_time.elapsed();
1013
1014 let start_decomp = std::time::Instant::now();
1015 if let Ok(decompressed) = compressor.decompress(&compressed) {
1016 let decompression_time = start_decomp.elapsed();
1017
1018 let error = calculate_compression_error(gradients, &decompressed);
1020
1021 let compression_ratio =
1022 compressed.data.len() as f64 / std::mem::size_of_val(gradients) as f64;
1023
1024 results.push(CompressionBenchmark {
1025 algorithm,
1026 compression_ratio,
1027 compression_time,
1028 decompression_time,
1029 error,
1030 compressed_size: compressed.data.len(),
1031 original_size: std::mem::size_of_val(gradients),
1032 });
1033 }
1034 }
1035 }
1036
1037 results
1038 }
1039
1040 fn calculate_compression_error<T: FloatElement + ToPrimitive>(
1042 original: &[T],
1043 reconstructed: &[T],
1044 ) -> f64 {
1045 if original.len() != reconstructed.len() {
1046 return f64::INFINITY;
1047 }
1048
1049 let mse = original
1050 .iter()
1051 .zip(reconstructed.iter())
1052 .map(|(&a, &b)| {
1053 let diff = ToPrimitive::to_f64(&a).expect("f64 conversion should succeed")
1054 - ToPrimitive::to_f64(&b).expect("f64 conversion should succeed");
1055 diff * diff
1056 })
1057 .sum::<f64>()
1058 / original.len() as f64;
1059
1060 mse
1061 }
1062}
1063
1064#[derive(Debug, Clone)]
1066pub struct GradientAnalysis {
1067 pub min_value: f64,
1068 pub max_value: f64,
1069 pub mean: f64,
1070 pub std_dev: f64,
1071 pub sparsity: f64,
1072 pub dynamic_range: f64,
1073 pub recommended_algorithm: CompressionAlgorithm,
1074}
1075
1076impl Default for GradientAnalysis {
1077 fn default() -> Self {
1078 Self {
1079 min_value: 0.0,
1080 max_value: 0.0,
1081 mean: 0.0,
1082 std_dev: 0.0,
1083 sparsity: 0.0,
1084 dynamic_range: 0.0,
1085 recommended_algorithm: CompressionAlgorithm::Quantization8Bit,
1086 }
1087 }
1088}
1089
1090#[derive(Debug, Clone)]
1092pub struct CompressionBenchmark {
1093 pub algorithm: CompressionAlgorithm,
1094 pub compression_ratio: f64,
1095 pub compression_time: std::time::Duration,
1096 pub decompression_time: std::time::Duration,
1097 pub error: f64,
1098 pub compressed_size: usize,
1099 pub original_size: usize,
1100}
1101
1102#[cfg(test)]
1103mod tests {
1104 use super::*;
1105 use approx::assert_relative_eq;
1106
1107 #[test]
1108 fn test_quantization_8bit() {
1109 let gradients: Vec<f32> = vec![0.1, 0.2, -0.3, 0.4, -0.5];
1110 let config = CompressionConfig {
1111 algorithm: CompressionAlgorithm::Quantization8Bit,
1112 ..Default::default()
1113 };
1114
1115 let mut compressor = GradientCompressor::new(config);
1116 let compressed = compressor.compress(&gradients, "test").unwrap();
1117 let decompressed = compressor.decompress(&compressed).unwrap();
1118
1119 assert_eq!(decompressed.len(), gradients.len());
1120
1121 for (i, (&original, &reconstructed)) in
1123 gradients.iter().zip(decompressed.iter()).enumerate()
1124 {
1125 let error = (original - reconstructed).abs();
1126 assert!(
1127 error < 0.01,
1128 "Value {} decompression error too large: {}",
1129 i,
1130 error
1131 );
1132 }
1133 }
1134
1135 #[test]
1136 fn test_top_k_sparsification() {
1137 let gradients: Vec<f32> = vec![0.1, 0.01, -0.3, 0.002, -0.5, 0.001];
1138 let config = CompressionConfig {
1139 algorithm: CompressionAlgorithm::TopKSparsification,
1140 sparsity_threshold: 0.5, ..Default::default()
1142 };
1143
1144 let mut compressor = GradientCompressor::new(config);
1145 let compressed = compressor.compress(&gradients, "test").unwrap();
1146 let decompressed = compressor.decompress(&compressed).unwrap();
1147
1148 assert_eq!(decompressed.len(), gradients.len());
1149
1150 assert_relative_eq!(decompressed[2], -0.3, epsilon = 0.1);
1152 assert_relative_eq!(decompressed[4], -0.5, epsilon = 0.1);
1153 }
1154
1155 #[test]
1156 fn test_adaptive_compression() {
1157 let gradients: Vec<f32> = vec![0.1, 0.0, 0.0, 0.2, 0.0, 0.0, 0.3]; let config = CompressionConfig {
1159 algorithm: CompressionAlgorithm::Adaptive,
1160 sparsity_threshold: 0.4, ..Default::default()
1162 };
1163
1164 let mut compressor = GradientCompressor::new(config);
1165 let compressed = compressor.compress(&gradients, "test").unwrap();
1166
1167 assert_eq!(
1169 compressed.algorithm,
1170 CompressionAlgorithm::TopKSparsification
1171 );
1172 }
1173
1174 #[test]
1175 fn test_compression_stats() {
1176 let gradients: Vec<f32> = vec![0.1, 0.2, 0.3, 0.4, 0.5];
1177 let config = CompressionConfig::default();
1178
1179 let mut compressor = GradientCompressor::new(config);
1180 compressor.compress(&gradients, "test").unwrap();
1181
1182 let stats = compressor.get_stats();
1183 assert_eq!(stats.total_compressions, 1);
1184 assert!(stats.total_bytes_original > 0);
1185 assert!(stats.total_bytes_compressed > 0);
1186 }
1187
1188 #[test]
1189 fn test_gradient_analysis() {
1190 let gradients: Vec<f32> = vec![0.1, 0.0, 0.2, 0.0, 0.3];
1191 let analysis = utils::analyze_gradients(&gradients);
1192
1193 assert_eq!(analysis.sparsity, 0.4); assert_relative_eq!(analysis.mean, 0.12, epsilon = 1e-6);
1195 assert!(analysis.std_dev > 0.0);
1196 }
1197}