//! Sparsification compression methods (TopK, Random, Threshold) use super::config::CompressionType; use super::types::CompressedGradient; use crate::error::Result; // ============================================================================= // TopK Sparsification // ============================================================================= pub fn compress_topk(gradients: &[f32], shape: &[usize], ratio: f32) -> CompressedGradient { let k = ((gradients.len() as f32 * ratio).ceil() as usize).max(1); // Find top-k absolute values let mut indexed: Vec<(usize, f32)> = gradients.iter().copied().enumerate().collect(); // Sort by absolute value descending indexed.sort_by(|a, b| b.1.abs().partial_cmp(&a.1.abs()).unwrap()); indexed.truncate(k); // Extract indices and values let indices: Vec = indexed.iter().map(|(i, _)| *i as u32).collect(); let values: Vec = indexed.iter().flat_map(|(_, v)| v.to_le_bytes()).collect(); CompressedGradient { shape: shape.to_vec(), compression_type: CompressionType::TopK, data: values, indices: Some(indices), scale: 1.0, zero_point: 0.0, num_elements: gradients.len(), original_size: gradients.len() * 4, } } // ============================================================================= // Random Sparsification // ============================================================================= pub fn compress_random(gradients: &[f32], shape: &[usize], ratio: f32) -> CompressedGradient { let k = ((gradients.len() as f32 * ratio).ceil() as usize).max(1); // Randomly select k indices use rand::seq::SliceRandom; let mut rng = rand::thread_rng(); let mut all_indices: Vec = (0..gradients.len()).collect(); all_indices.shuffle(&mut rng); all_indices.truncate(k); all_indices.sort_unstable(); let indices: Vec = all_indices.iter().map(|&i| i as u32).collect(); let values: Vec = all_indices .iter() .flat_map(|&i| gradients[i].to_le_bytes()) .collect(); // Scale values to compensate for sparsification let scale = 1.0 / ratio; CompressedGradient { shape: shape.to_vec(), compression_type: CompressionType::RandomK, data: values, indices: Some(indices), scale, zero_point: 0.0, num_elements: gradients.len(), original_size: gradients.len() * 4, } } // ============================================================================= // Threshold Sparsification // ============================================================================= pub fn compress_threshold( gradients: &[f32], shape: &[usize], threshold: f32, ) -> CompressedGradient { // Keep only values above threshold let mut indices: Vec = Vec::new(); let mut values: Vec = Vec::new(); for (i, &g) in gradients.iter().enumerate() { if g.abs() > threshold { indices.push(i as u32); values.extend_from_slice(&g.to_le_bytes()); } } CompressedGradient { shape: shape.to_vec(), compression_type: CompressionType::Threshold, data: values, indices: Some(indices), scale: 1.0, zero_point: 0.0, num_elements: gradients.len(), original_size: gradients.len() * 4, } } // ============================================================================= // Sparse Decompression (shared by TopK, Random, Threshold, DGC) // ============================================================================= pub fn decompress_sparse(compressed: &CompressedGradient) -> Result> { let mut result = vec![0.0f32; compressed.num_elements]; if let Some(indices) = &compressed.indices { let values: Vec = compressed .data .chunks_exact(4) .map(|chunk| { let bytes: [u8; 4] = chunk.try_into().unwrap(); f32::from_le_bytes(bytes) }) .collect(); for (&idx, &val) in indices.iter().zip(values.iter()) { result[idx as usize] = val * compressed.scale; } } Ok(result) }