//! Magnitude-based pruning implementation use crate::error::{CompressionError, Result}; use crate::pruning::config::{PruningConfig, PruningGranularity, PruningStatistics}; use rtx_tensor::{Device, Tensor}; /// Magnitude-based pruner for unstructured pruning #[derive(Debug, Clone)] pub struct MagnitudePruner { config: PruningConfig, device: Device, schedule_steps: Option, initial_sparsity: f32, final_sparsity: f32, n_nonzeros: Option, m_blocksize: Option, } impl MagnitudePruner { /// Create a new magnitude pruner pub fn new(config: PruningConfig, device: &Device) -> Result { if config.sparsity < 0.0 || config.sparsity > 1.0 { return Err(CompressionError::CompressionFailed( "Sparsity must be between 0 and 1".to_string(), )); } let final_sparsity = config.sparsity; Ok(Self { config, device: device.clone(), schedule_steps: None, initial_sparsity: 0.0, final_sparsity, n_nonzeros: None, m_blocksize: None, }) } /// Create pruner with gradual sparsity schedule pub fn with_schedule( final_sparsity: f32, schedule_steps: usize, device: &Device, ) -> Result { let config = PruningConfig { sparsity: final_sparsity, structured: false, granularity: PruningGranularity::Unstructured, preserve_gradients: false, }; Ok(Self { config, device: device.clone(), schedule_steps: Some(schedule_steps), initial_sparsity: 0.0, final_sparsity, n_nonzeros: None, m_blocksize: None, }) } /// Create pruner with N:M sparsity pattern pub fn with_nm_sparsity(n: usize, m: usize, device: &Device) -> Result { if n >= m { return Err(CompressionError::CompressionFailed( "N must be less than M for N:M sparsity".to_string(), )); } let sparsity = 1.0 - (n as f32 / m as f32); let config = PruningConfig { sparsity, structured: false, granularity: PruningGranularity::Block, preserve_gradients: false, }; Ok(Self { config, device: device.clone(), schedule_steps: None, initial_sparsity: 0.0, final_sparsity: sparsity, n_nonzeros: Some(n), m_blocksize: Some(m), }) } /// Compute pruning mask for weights pub fn compute_mask(&self, weights: &Tensor) -> Result { let data = weights.to_vec()?; let num_elements = data.len(); let num_to_prune = (num_elements as f32 * self.config.sparsity) as usize; // Compute magnitudes let mut magnitudes: Vec<(f32, usize)> = data .iter() .enumerate() .map(|(i, &v)| (v.abs(), i)) .collect(); // Sort by magnitude magnitudes.sort_by(|a, b| a.0.total_cmp(&b.0)); // Create mask (1 for keep, 0 for prune) let mut mask = vec![1.0f32; num_elements]; for i in 0..num_to_prune { mask[magnitudes[i].1] = 0.0; } Ok(Tensor::from_data( mask, weights.shape().dims().to_vec(), &self.device, )?) } /// Compute mask at specific training step (for gradual pruning) pub fn compute_mask_at_step(&self, weights: &Tensor, step: usize) -> Result { let current_sparsity = self.get_current_sparsity(step); // Temporarily update config let mut temp_pruner = self.clone(); temp_pruner.config.sparsity = current_sparsity; temp_pruner.compute_mask(weights) } /// Compute mask with importance scores pub fn compute_mask_with_importance( &self, weights: &Tensor, importance: &Tensor, ) -> Result { if weights.shape() != importance.shape() { return Err(CompressionError::CompressionFailed( "Weights and importance must have same shape".to_string(), )); } let importance_data = importance.to_vec()?; let num_elements = importance_data.len(); let num_to_prune = (num_elements as f32 * self.config.sparsity) as usize; // Sort by importance (lower importance = prune first) let mut scores: Vec<(f32, usize)> = importance_data .iter() .enumerate() .map(|(i, &v)| (v, i)) .collect(); scores.sort_by(|a, b| a.0.total_cmp(&b.0)); // Create mask let mut mask = vec![1.0f32; num_elements]; for i in 0..num_to_prune { mask[scores[i].1] = 0.0; } Ok(Tensor::from_data( mask, weights.shape().dims().to_vec(), &self.device, )?) } /// Apply pruning mask to weights pub fn apply_mask(&self, weights: &Tensor, mask: &Tensor) -> Result { Ok(weights.mul(mask)?) } /// Get current sparsity at training step pub fn get_current_sparsity(&self, step: usize) -> f32 { if let Some(total_steps) = self.schedule_steps { let progress = (step as f32) / (total_steps as f32).max(1.0); let progress = progress.min(1.0); // Cubic schedule for gradual pruning let t = progress.powi(3); self.initial_sparsity + (self.final_sparsity - self.initial_sparsity) * t } else { self.config.sparsity } } /// Analyze weights and return statistics pub fn analyze(&self, weights: &Tensor) -> Result { let data = weights.to_vec()?; let total = data.len(); let zeros = data.iter().filter(|&&x| x == 0.0).count(); Ok(PruningStatistics { total_parameters: total, pruned_parameters: zeros, target_sparsity: self.config.sparsity, achieved_sparsity: zeros as f32 / total as f32, compression_ratio: if zeros > 0 { total as f32 / (total - zeros) as f32 } else { 1.0 }, }) } // Getters for testing pub fn sparsity(&self) -> f32 { self.config.sparsity } pub fn is_structured(&self) -> bool { self.config.structured } }