Files
rustytorch/crates/training/rtx-distributed/src/gradient_compression/sparsification.rs
T
2026-03-04 00:08:42 +00:00

129 lines
4.1 KiB
Rust

//! 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<u32> = indexed.iter().map(|(i, _)| *i as u32).collect();
let values: Vec<u8> = 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<usize> = (0..gradients.len()).collect();
all_indices.shuffle(&mut rng);
all_indices.truncate(k);
all_indices.sort_unstable();
let indices: Vec<u32> = all_indices.iter().map(|&i| i as u32).collect();
let values: Vec<u8> = 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<u32> = Vec::new();
let mut values: Vec<u8> = 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<Vec<f32>> {
let mut result = vec![0.0f32; compressed.num_elements];
if let Some(indices) = &compressed.indices {
let values: Vec<f32> = 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)
}