use crate::Result;
use super::{WeightMatrix, NormalizationMetadata};
#[derive(Debug)]
pub struct DistributionNormalizer {
outlier_threshold: f32,
seed: u64,
}
impl DistributionNormalizer {
pub fn new(outlier_threshold: f32, seed: u64) -> Self {
Self {
outlier_threshold,
seed,
}
}
pub fn normalize(&self, weights: &mut WeightMatrix) -> Result<NormalizationMetadata> {
let rows = weights.rows();
let cols = weights.cols();
let mut channel_means = Vec::with_capacity(rows);
let mut channel_variances = Vec::with_capacity(rows);
for i in 0..rows {
let row = weights.get_row(i);
let mean = row.iter().sum::<f32>() / cols as f32;
channel_means.push(mean);
let variance = row.iter()
.map(|&x| (x - mean).powi(2))
.sum::<f32>() / cols as f32;
channel_variances.push(variance);
}
let outlier_channels = self.identify_outlier_channels(&channel_variances);
let mut channel_scales = vec![1.0; rows];
let delta = self.calculate_delta(&channel_variances);
for &channel_idx in &outlier_channels {
let sigma = channel_variances[channel_idx].sqrt();
channel_scales[channel_idx] = sigma / delta;
}
for i in 0..rows {
let mean = channel_means[i];
let scale = channel_scales[i];
let row = weights.get_row_mut(i);
for j in 0..cols {
row[j] = (row[j] - mean) / scale;
}
}
Ok(NormalizationMetadata {
channel_means,
channel_scales,
outlier_channels,
})
}
pub fn denormalize(&self, weights: &mut [f32], metadata: &NormalizationMetadata) -> Result<()> {
let rows = metadata.channel_means.len();
let cols = weights.len() / rows;
for i in 0..rows {
let mean = metadata.channel_means[i];
let scale = metadata.channel_scales[i];
let start = i * cols;
let end = start + cols;
for j in start..end {
weights[j] = weights[j] * scale + mean;
}
}
Ok(())
}
fn identify_outlier_channels(&self, variances: &[f32]) -> Vec<usize> {
let mut indexed_variances: Vec<(usize, f32)> = variances.iter()
.enumerate()
.map(|(i, &v)| (i, v))
.collect();
indexed_variances.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
let num_outliers = ((variances.len() as f32) * self.outlier_threshold).ceil() as usize;
let num_outliers = num_outliers.max(1);
indexed_variances.into_iter()
.take(num_outliers)
.map(|(idx, _)| idx)
.collect()
}
fn calculate_delta(&self, variances: &[f32]) -> f32 {
let mut sorted_variances = variances.to_vec();
sorted_variances.sort_by(|a, b| a.partial_cmp(b).unwrap());
let median_idx = sorted_variances.len() / 2;
let median_variance = if sorted_variances.len() % 2 == 0 {
(sorted_variances[median_idx - 1] + sorted_variances[median_idx]) / 2.0
} else {
sorted_variances[median_idx]
};
median_variance.sqrt().max(1e-8) }
}
pub struct ChannelStatistics {
pub means: Vec<f32>,
pub variances: Vec<f32>,
pub std_devs: Vec<f32>,
pub outlier_ratio: f32,
}
impl ChannelStatistics {
pub fn compute(weights: &WeightMatrix, outlier_threshold: f32) -> Self {
let rows = weights.rows();
let cols = weights.cols();
let mut means = Vec::with_capacity(rows);
let mut variances = Vec::with_capacity(rows);
let mut std_devs = Vec::with_capacity(rows);
for i in 0..rows {
let row = weights.get_row(i);
let mean = row.iter().sum::<f32>() / cols as f32;
means.push(mean);
let variance = row.iter()
.map(|&x| (x - mean).powi(2))
.sum::<f32>() / cols as f32;
variances.push(variance);
std_devs.push(variance.sqrt());
}
Self {
means,
variances,
std_devs,
outlier_ratio: outlier_threshold,
}
}
pub fn print_summary(&self) {
println!("Channel Statistics Summary:");
println!(" Total channels: {}", self.means.len());
println!(" Mean of means: {:.6}", self.means.iter().sum::<f32>() / self.means.len() as f32);
println!(" Mean variance: {:.6}", self.variances.iter().sum::<f32>() / self.variances.len() as f32);
println!(" Max variance: {:.6}", self.variances.iter().fold(0.0f32, |a, &b| a.max(b)));
println!(" Min variance: {:.6}", self.variances.iter().fold(f32::INFINITY, |a, &b| a.min(b)));
println!(" Outlier threshold: {:.1}%", self.outlier_ratio * 100.0);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_normalization_basic() {
let data = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, ];
let mut weights = WeightMatrix::new(data, vec![2, 3], "test".to_string());
let normalizer = DistributionNormalizer::new(0.5, 42);
let metadata = normalizer.normalize(&mut weights).unwrap();
assert_eq!(metadata.channel_means.len(), 2);
assert_eq!(metadata.channel_scales.len(), 2);
assert!((metadata.channel_means[0] - 2.0).abs() < 1e-6);
assert!((metadata.channel_means[1] - 5.0).abs() < 1e-6);
}
#[test]
fn test_outlier_identification() {
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);
assert_eq!(outliers.len(), 2);
assert!(outliers.contains(&1)); assert!(outliers.contains(&4)); }
#[test]
fn test_denormalization() {
let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut weights = WeightMatrix::new(data.clone(), vec![2, 3], "test".to_string());
let normalizer = DistributionNormalizer::new(0.5, 42);
let metadata = normalizer.normalize(&mut weights).unwrap();
let mut denormalized = weights.data.clone();
normalizer.denormalize(&mut denormalized, &metadata).unwrap();
for (orig, denorm) in data.iter().zip(denormalized.iter()) {
assert!((orig - denorm).abs() < 1e-5, "Original: {}, Denormalized: {}", orig, denorm);
}
}
#[test]
fn test_delta_calculation() {
let variances = vec![1.0, 4.0, 9.0, 16.0, 25.0]; let normalizer = DistributionNormalizer::new(0.2, 42);
let delta = normalizer.calculate_delta(&variances);
assert!((delta - 3.0).abs() < 1e-6);
}
}