Skip to main content

ohms_adaptq/novaq/
normalization.rs

1use crate::Result;
2use super::{WeightMatrix, NormalizationMetadata};
3
4/// Distribution Normalizer implementing Stage 1 of NOVAQ
5/// 
6/// Mathematical formulation from NOVAQ paper:
7/// W_hat_{i,:} = (W_{i,:} - μ_i) / s_i
8/// 
9/// where:
10/// μ_i = (1/d) * Σ_j W_{i,j}  (per-channel mean)
11/// s_i = σ_i / Δ if σ_i in top p%, else 1  (outlier scaling)
12#[derive(Debug)]
13pub struct DistributionNormalizer {
14    outlier_threshold: f32,
15    seed: u64,
16}
17
18impl DistributionNormalizer {
19    pub fn new(outlier_threshold: f32, seed: u64) -> Self {
20        Self {
21            outlier_threshold,
22            seed,
23        }
24    }
25    
26    /// Normalize weight matrix by eliminating per-channel means and rescaling outlier channels
27    pub fn normalize(&self, weights: &mut WeightMatrix) -> Result<NormalizationMetadata> {
28        let rows = weights.rows();
29        let cols = weights.cols();
30        
31        let mut channel_means = Vec::with_capacity(rows);
32        let mut channel_variances = Vec::with_capacity(rows);
33        
34        // Calculate per-channel statistics
35        for i in 0..rows {
36            let row = weights.get_row(i);
37            
38            // Calculate mean: μ_i = (1/d) * Σ_j W_{i,j}
39            let mean = row.iter().sum::<f32>() / cols as f32;
40            channel_means.push(mean);
41            
42            // Calculate variance for outlier detection
43            let variance = row.iter()
44                .map(|&x| (x - mean).powi(2))
45                .sum::<f32>() / cols as f32;
46            channel_variances.push(variance);
47        }
48        
49        // Identify outlier channels (top-p% by variance)
50        let outlier_channels = self.identify_outlier_channels(&channel_variances);
51        
52        // Calculate scaling factors
53        let mut channel_scales = vec![1.0; rows];
54        let delta = self.calculate_delta(&channel_variances);
55        
56        for &channel_idx in &outlier_channels {
57            let sigma = channel_variances[channel_idx].sqrt();
58            channel_scales[channel_idx] = sigma / delta;
59        }
60        
61        // Apply normalization: W_hat_{i,:} = (W_{i,:} - μ_i) / s_i
62        for i in 0..rows {
63            let mean = channel_means[i];
64            let scale = channel_scales[i];
65            let row = weights.get_row_mut(i);
66            
67            for j in 0..cols {
68                row[j] = (row[j] - mean) / scale;
69            }
70        }
71        
72        Ok(NormalizationMetadata {
73            channel_means,
74            channel_scales,
75            outlier_channels,
76        })
77    }
78    
79    /// Denormalize weights using stored metadata
80    pub fn denormalize(&self, weights: &mut [f32], metadata: &NormalizationMetadata) -> Result<()> {
81        let rows = metadata.channel_means.len();
82        let cols = weights.len() / rows;
83        
84        for i in 0..rows {
85            let mean = metadata.channel_means[i];
86            let scale = metadata.channel_scales[i];
87            
88            let start = i * cols;
89            let end = start + cols;
90            
91            for j in start..end {
92                weights[j] = weights[j] * scale + mean;
93            }
94        }
95        
96        Ok(())
97    }
98    
99    /// Identify channels with top-p% variance as outliers
100    fn identify_outlier_channels(&self, variances: &[f32]) -> Vec<usize> {
101        let mut indexed_variances: Vec<(usize, f32)> = variances.iter()
102            .enumerate()
103            .map(|(i, &v)| (i, v))
104            .collect();
105        
106        // Sort by variance in descending order
107        indexed_variances.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
108        
109        // Take top-p% channels
110        let num_outliers = ((variances.len() as f32) * self.outlier_threshold).ceil() as usize;
111        let num_outliers = num_outliers.max(1); // At least one outlier
112        
113        indexed_variances.into_iter()
114            .take(num_outliers)
115            .map(|(idx, _)| idx)
116            .collect()
117    }
118    
119    /// Calculate delta parameter for outlier scaling
120    /// Uses median variance as a robust estimator
121    fn calculate_delta(&self, variances: &[f32]) -> f32 {
122        let mut sorted_variances = variances.to_vec();
123        sorted_variances.sort_by(|a, b| a.partial_cmp(b).unwrap());
124        
125        let median_idx = sorted_variances.len() / 2;
126        let median_variance = if sorted_variances.len() % 2 == 0 {
127            (sorted_variances[median_idx - 1] + sorted_variances[median_idx]) / 2.0
128        } else {
129            sorted_variances[median_idx]
130        };
131        
132        median_variance.sqrt().max(1e-8) // Avoid division by zero
133    }
134}
135
136/// Calculate channel-wise statistics for analysis
137pub struct ChannelStatistics {
138    pub means: Vec<f32>,
139    pub variances: Vec<f32>,
140    pub std_devs: Vec<f32>,
141    pub outlier_ratio: f32,
142}
143
144impl ChannelStatistics {
145    pub fn compute(weights: &WeightMatrix, outlier_threshold: f32) -> Self {
146        let rows = weights.rows();
147        let cols = weights.cols();
148        
149        let mut means = Vec::with_capacity(rows);
150        let mut variances = Vec::with_capacity(rows);
151        let mut std_devs = Vec::with_capacity(rows);
152        
153        for i in 0..rows {
154            let row = weights.get_row(i);
155            
156            let mean = row.iter().sum::<f32>() / cols as f32;
157            means.push(mean);
158            
159            let variance = row.iter()
160                .map(|&x| (x - mean).powi(2))
161                .sum::<f32>() / cols as f32;
162            variances.push(variance);
163            std_devs.push(variance.sqrt());
164        }
165        
166        Self {
167            means,
168            variances,
169            std_devs,
170            outlier_ratio: outlier_threshold,
171        }
172    }
173    
174    pub fn print_summary(&self) {
175        println!("Channel Statistics Summary:");
176        println!("  Total channels: {}", self.means.len());
177        println!("  Mean of means: {:.6}", self.means.iter().sum::<f32>() / self.means.len() as f32);
178        println!("  Mean variance: {:.6}", self.variances.iter().sum::<f32>() / self.variances.len() as f32);
179        println!("  Max variance: {:.6}", self.variances.iter().fold(0.0f32, |a, &b| a.max(b)));
180        println!("  Min variance: {:.6}", self.variances.iter().fold(f32::INFINITY, |a, &b| a.min(b)));
181        println!("  Outlier threshold: {:.1}%", self.outlier_ratio * 100.0);
182    }
183}
184
185#[cfg(test)]
186mod tests {
187    use super::*;
188    
189    #[test]
190    fn test_normalization_basic() {
191        let data = vec![
192            1.0, 2.0, 3.0,  // Row 0: mean = 2.0
193            4.0, 5.0, 6.0,  // Row 1: mean = 5.0
194        ];
195        let mut weights = WeightMatrix::new(data, vec![2, 3], "test".to_string());
196        
197        let normalizer = DistributionNormalizer::new(0.5, 42);
198        let metadata = normalizer.normalize(&mut weights).unwrap();
199        
200        assert_eq!(metadata.channel_means.len(), 2);
201        assert_eq!(metadata.channel_scales.len(), 2);
202        assert!((metadata.channel_means[0] - 2.0).abs() < 1e-6);
203        assert!((metadata.channel_means[1] - 5.0).abs() < 1e-6);
204    }
205    
206    #[test]
207    fn test_outlier_identification() {
208        let variances = vec![1.0, 100.0, 1.5, 2.0, 150.0]; // Indices 1 and 4 are outliers
209        let normalizer = DistributionNormalizer::new(0.4, 42); // Top 40% = 2 channels
210        let outliers = normalizer.identify_outlier_channels(&variances);
211        
212        assert_eq!(outliers.len(), 2);
213        assert!(outliers.contains(&1)); // Variance 100.0
214        assert!(outliers.contains(&4)); // Variance 150.0
215    }
216    
217    #[test]
218    fn test_denormalization() {
219        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
220        let mut weights = WeightMatrix::new(data.clone(), vec![2, 3], "test".to_string());
221        
222        let normalizer = DistributionNormalizer::new(0.5, 42);
223        let metadata = normalizer.normalize(&mut weights).unwrap();
224        
225        // Denormalize
226        let mut denormalized = weights.data.clone();
227        normalizer.denormalize(&mut denormalized, &metadata).unwrap();
228        
229        // Should be close to original
230        for (orig, denorm) in data.iter().zip(denormalized.iter()) {
231            assert!((orig - denorm).abs() < 1e-5, "Original: {}, Denormalized: {}", orig, denorm);
232        }
233    }
234    
235    #[test]
236    fn test_delta_calculation() {
237        let variances = vec![1.0, 4.0, 9.0, 16.0, 25.0]; // Medians: 9.0, sqrt = 3.0
238        let normalizer = DistributionNormalizer::new(0.2, 42);
239        let delta = normalizer.calculate_delta(&variances);
240        
241        assert!((delta - 3.0).abs() < 1e-6);
242    }
243}