ohms_adaptq/novaq/
normalization.rs1use crate::Result;
2use super::{WeightMatrix, NormalizationMetadata};
3
4#[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 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 for i in 0..rows {
36 let row = weights.get_row(i);
37
38 let mean = row.iter().sum::<f32>() / cols as f32;
40 channel_means.push(mean);
41
42 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 let outlier_channels = self.identify_outlier_channels(&channel_variances);
51
52 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 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 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 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 indexed_variances.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
108
109 let num_outliers = ((variances.len() as f32) * self.outlier_threshold).ceil() as usize;
111 let num_outliers = num_outliers.max(1); indexed_variances.into_iter()
114 .take(num_outliers)
115 .map(|(idx, _)| idx)
116 .collect()
117 }
118
119 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) }
134}
135
136pub 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, 4.0, 5.0, 6.0, ];
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]; let normalizer = DistributionNormalizer::new(0.4, 42); let outliers = normalizer.identify_outlier_channels(&variances);
211
212 assert_eq!(outliers.len(), 2);
213 assert!(outliers.contains(&1)); assert!(outliers.contains(&4)); }
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 let mut denormalized = weights.data.clone();
227 normalizer.denormalize(&mut denormalized, &metadata).unwrap();
228
229 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]; 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}