Skip to main content

torsh_autograd/
compression.rs

1//! Advanced gradient compression techniques for memory-limited environments
2//!
3//! This module provides sophisticated compression algorithms for gradients that can
4//! significantly reduce memory usage and communication overhead in distributed training.
5
6use 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/// Gradient compression configuration
15#[derive(Debug, Clone)]
16pub struct CompressionConfig {
17    /// Compression algorithm to use
18    pub algorithm: CompressionAlgorithm,
19    /// Target compression ratio (0.0 to 1.0)
20    pub target_ratio: f64,
21    /// Error tolerance for lossy compression
22    pub error_tolerance: f64,
23    /// Memory budget in bytes
24    pub memory_budget: usize,
25    /// Whether to use adaptive compression
26    pub adaptive: bool,
27    /// Minimum sparsity threshold for sparse compression
28    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, // 4x compression
36            error_tolerance: 1e-4,
37            memory_budget: 100 * 1024 * 1024, // 100MB
38            adaptive: true,
39            sparsity_threshold: 0.01, // 1% sparsity threshold
40        }
41    }
42}
43
44/// Available compression algorithms
45#[derive(Debug, Clone, Copy, PartialEq, Eq)]
46pub enum CompressionAlgorithm {
47    /// No compression
48    None,
49    /// 8-bit quantization
50    Quantization8Bit,
51    /// 4-bit quantization
52    Quantization4Bit,
53    /// 2-bit quantization
54    Quantization2Bit,
55    /// 1-bit quantization (sign only)
56    Quantization1Bit,
57    /// Top-K sparsification
58    TopKSparsification,
59    /// Random sparsification
60    RandomSparsification,
61    /// Gradient sketching with random projections
62    GradientSketching,
63    /// PowerSGD low-rank compression
64    PowerSGD,
65    /// Error feedback compression
66    ErrorFeedback,
67    /// Adaptive compression (chooses best algorithm)
68    Adaptive,
69}
70
71/// Compressed gradient representation
72#[derive(Debug, Clone)]
73pub struct CompressedGradient {
74    /// Original shape of the gradient
75    pub original_shape: Vec<usize>,
76    /// Compressed data
77    pub data: Vec<u8>,
78    /// Compression metadata
79    pub metadata: CompressionMetadata,
80    /// Algorithm used for compression
81    pub algorithm: CompressionAlgorithm,
82}
83
84/// Compression metadata
85#[derive(Debug, Clone)]
86pub struct CompressionMetadata {
87    /// Scale factor for quantization
88    pub scale: f64,
89    /// Zero point for quantization
90    pub zero_point: i32,
91    /// Indices for sparse compression
92    pub indices: Vec<usize>,
93    /// Random seed for reproducible compression
94    pub seed: u64,
95    /// Rank for low-rank compression
96    pub rank: usize,
97    /// Error compensation
98    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/// Gradient compressor with advanced algorithms
115#[derive(Clone)]
116pub struct GradientCompressor<T: FloatElement> {
117    /// Compression configuration
118    config: CompressionConfig,
119    /// Error feedback buffer for each parameter
120    error_feedback: Arc<RwLock<HashMap<String, Vec<T>>>>,
121    /// Compression statistics
122    stats: Arc<RwLock<CompressionStats>>,
123    /// Random number generator state
124    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/// Compression statistics
141#[derive(Debug, Clone, Default)]
142pub struct CompressionStats {
143    /// Total number of compressions
144    pub total_compressions: usize,
145    /// Total bytes before compression
146    pub total_bytes_original: usize,
147    /// Total bytes after compression
148    pub total_bytes_compressed: usize,
149    /// Average compression ratio
150    pub avg_compression_ratio: f64,
151    /// Total compression time
152    pub total_compression_time_ms: u64,
153    /// Total decompression time
154    pub total_decompression_time_ms: u64,
155    /// Compression error (if applicable)
156    pub compression_error: f64,
157}
158
159impl<T: FloatElement + FromPrimitive + ToPrimitive> GradientCompressor<T> {
160    /// Create a new gradient compressor
161    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)), // Fixed seed for reproducibility
167        }
168    }
169
170    /// Compress gradients according to the configured algorithm
171    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        // Update statistics
211        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    /// Decompress gradients
229    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        // Update decompression statistics
265        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    /// Choose the best compression algorithm based on gradient characteristics
272    fn choose_best_algorithm(&self, gradients: &[T]) -> Result<CompressionAlgorithm> {
273        // Analyze gradient characteristics
274        let sparsity = self.calculate_sparsity(gradients);
275        let variance = self.calculate_variance(gradients);
276        let magnitude = self.calculate_magnitude(gradients);
277
278        // Choose algorithm based on characteristics
279        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    /// Calculate sparsity (fraction of near-zero elements)
293    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    /// Calculate variance of gradients
301    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    /// Calculate magnitude (RMS) of gradients
325    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    /// No compression (passthrough)
343    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    /// 8-bit quantization compression
361    fn compress_quantization_8bit(&self, gradients: &[T]) -> Result<CompressedGradient> {
362        // Find min and max values
363        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        // Calculate scale and zero point
373        let scale = (max_val - min_val) / 255.0;
374        let zero_point = (-min_val / scale).round() as i32;
375
376        // Quantize gradients
377        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    /// 4-bit quantization compression
399    fn compress_quantization_4bit(&self, gradients: &[T]) -> Result<CompressedGradient> {
400        // Similar to 8-bit but with 4-bit precision
401        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; // 4-bit has 16 levels (0-15)
411        let zero_point = (-min_val / scale).round() as i32;
412
413        // Pack two 4-bit values into each byte
414        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    /// 2-bit quantization compression
448    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; // 2-bit has 4 levels (0-3)
459        let zero_point = (-min_val / scale).round() as i32;
460
461        // Pack four 2-bit values into each byte
462        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    /// 1-bit quantization (sign only)
490    fn compress_quantization_1bit(&self, gradients: &[T]) -> Result<CompressedGradient> {
491        // Calculate magnitude for scaling
492        let magnitude = self.calculate_magnitude(gradients);
493
494        // Pack eight 1-bit values into each byte
495        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    /// Top-K sparsification
521    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        // Create (value, index) pairs and sort by absolute value
525        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        // Keep only top-k values
542        let top_k_indices: Vec<usize> = value_index_pairs
543            .iter()
544            .take(k)
545            .map(|(_, idx)| *idx)
546            .collect();
547
548        // Store compressed values and indices
549        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    /// Random sparsification
573    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            // Simple linear congruential generator
582            *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    /// Gradient sketching compression
608    fn compress_gradient_sketching(&self, gradients: &[T]) -> Result<CompressedGradient> {
609        // Reduce dimensionality using random projections
610        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            // Apply random projections
620            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        // Convert sketch to bytes
629        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    /// PowerSGD low-rank compression
650    fn compress_power_sgd(&self, gradients: &[T]) -> Result<CompressedGradient> {
651        // For simplicity, we'll use a basic low-rank approximation
652        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        // Create a simplified low-rank representation
658        // In a full implementation, this would use SVD or power iteration
659        let mut compressed_data = Vec::new();
660
661        // Store first 'rank' values as the low-rank approximation
662        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    /// Error feedback compression
682    fn compress_error_feedback(
683        &mut self,
684        gradients: &[T],
685        parameter_name: &str,
686    ) -> Result<CompressedGradient> {
687        // Get or create error feedback buffer
688        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        // Ensure buffer has correct size
696        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        // Add error feedback to gradients
704        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        // Compress using quantization
710        let compressed = self.compress_quantization_8bit(&compensated_gradients)?;
711
712        // Calculate new error
713        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    // Decompression methods (implementations for each algorithm)
730    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            // Extract first 4-bit value
765            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                // Extract second 4-bit value
774                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        // Gradient sketching is lossy and cannot be perfectly reconstructed
892        // This would require additional information or approximation
893        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        // Simple reconstruction: broadcast first 'rank' values
905        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        // Error feedback uses the same decompression as the underlying algorithm
923        self.decompress_quantization_8bit(compressed)
924    }
925
926    /// Get compression statistics
927    pub fn get_stats(&self) -> CompressionStats {
928        self.stats.read_or_recover().clone()
929    }
930
931    /// Reset compression statistics
932    pub fn reset_stats(&mut self) {
933        *self.stats.write_or_recover() = CompressionStats::default();
934    }
935
936    /// Update configuration
937    pub fn update_config(&mut self, new_config: CompressionConfig) {
938        self.config = new_config;
939    }
940}
941
942/// Utilities for compression analysis
943pub mod utils {
944    use super::*;
945
946    /// Analyze gradient characteristics to suggest optimal compression
947    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    /// Benchmark different compression algorithms
996    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                    // Calculate error
1019                    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    /// Calculate compression error (MSE)
1041    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/// Gradient analysis results
1065#[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/// Compression benchmark results
1091#[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        // Check that decompressed values are close to original (within quantization error)
1122        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, // Keep 50%
1141            ..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        // Check that large values are preserved (with some tolerance for compression)
1151        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]; // Sparse
1158        let config = CompressionConfig {
1159            algorithm: CompressionAlgorithm::Adaptive,
1160            sparsity_threshold: 0.4, // 40% sparsity threshold
1161            ..Default::default()
1162        };
1163
1164        let mut compressor = GradientCompressor::new(config);
1165        let compressed = compressor.compress(&gradients, "test").unwrap();
1166
1167        // Should choose sparsification for sparse data
1168        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); // 2 out of 5 are zero
1194        assert_relative_eq!(analysis.mean, 0.12, epsilon = 1e-6);
1195        assert!(analysis.std_dev > 0.0);
1196    }
1197}