129 lines
4.1 KiB
Rust
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)
|
|
}
|