//! Aggregation strategies for federated learning. //! //! This module implements various aggregation algorithms for combining //! client model updates in a federated learning setting. use fedmed_shared::{AggregationStrategy, GlobalModel, LocalUpdate}; use crate::FederatedError; /// Aggregator for combining client updates. #[derive(Debug)] pub struct Aggregator { /// Aggregation strategy. strategy: AggregationStrategy, /// FedProx mu parameter. fedprox_mu: f64, /// Scaffold server learning rate. scaffold_server_lr: f64, /// Control variates for Scaffold. control_variate: Option>, /// Trimmed mean fraction (for robust aggregation). trim_fraction: f32, } impl Default for Aggregator { fn default() -> Self { Self::new(AggregationStrategy::FedAvg) } } impl Aggregator { /// Create a new aggregator with the specified strategy. #[must_use] pub fn new(strategy: AggregationStrategy) -> Self { Self { strategy, fedprox_mu: 0.01, scaffold_server_lr: 1.0, control_variate: None, trim_fraction: 0.1, } } /// Create a FedProx aggregator. #[must_use] pub fn fedprox(mu: f64) -> Self { Self { strategy: AggregationStrategy::FedProx, fedprox_mu: mu, scaffold_server_lr: 1.0, control_variate: None, trim_fraction: 0.1, } } /// Create a Scaffold aggregator. #[must_use] pub fn scaffold(server_lr: f64) -> Self { Self { strategy: AggregationStrategy::Scaffold, fedprox_mu: 0.01, scaffold_server_lr: server_lr, control_variate: None, trim_fraction: 0.1, } } /// Aggregate client updates. pub fn aggregate( &mut self, updates: &mut [LocalUpdate], global_model: &GlobalModel, ) -> Result, FederatedError> { if updates.is_empty() { return Err(FederatedError::AggregationFailed( "No updates to aggregate".to_string(), )); } match self.strategy { AggregationStrategy::FedAvg => self.fedavg(updates), AggregationStrategy::FedProx => self.fedprox_aggregate(updates, global_model), AggregationStrategy::Scaffold => self.scaffold_aggregate(updates, global_model), AggregationStrategy::FedAdam => self.fedadam(updates), AggregationStrategy::FedYogi => self.fedyogi(updates), AggregationStrategy::TrimmedMean => self.trimmed_mean(updates), AggregationStrategy::Median => self.median_aggregate(updates), AggregationStrategy::Krum => self.krum(updates), } } /// FedAvg: Weighted averaging of model updates. fn fedavg(&self, updates: &[LocalUpdate]) -> Result, FederatedError> { if updates.is_empty() { return Err(FederatedError::AggregationFailed( "No updates for FedAvg".to_string(), )); } let total_samples: usize = updates.iter().map(|u| u.num_samples).sum(); if total_samples == 0 { return Err(FederatedError::AggregationFailed( "Total samples is zero".to_string(), )); } let num_params = updates[0].weight_deltas.len(); let mut aggregated = vec![0.0_f32; num_params]; for update in updates { let weight = update.num_samples as f32 / total_samples as f32; for (i, delta) in update.weight_deltas.iter().enumerate() { if i < aggregated.len() { aggregated[i] += delta * weight; } } } Ok(aggregated) } /// FedProx: Adds proximal term to handle heterogeneous data. fn fedprox_aggregate( &self, updates: &[LocalUpdate], global_model: &GlobalModel, ) -> Result, FederatedError> { // First do FedAvg let mut aggregated = self.fedavg(updates)?; // Apply proximal regularization // The proximal term: mu/2 * ||w - w_global||^2 // This is typically applied during local training, but we can // approximate the effect here by damping large updates let mu = self.fedprox_mu as f32; for (i, delta) in aggregated.iter_mut().enumerate() { if i < global_model.weights.len() { // Dampen updates that are far from the global model let damping = 1.0 / (1.0 + mu * delta.abs()); *delta *= damping; } } Ok(aggregated) } /// Scaffold: Uses control variates to correct client drift. fn scaffold_aggregate( &mut self, updates: &mut [LocalUpdate], global_model: &GlobalModel, ) -> Result, FederatedError> { let num_params = updates[0].weight_deltas.len(); // Initialize control variate if needed if self.control_variate.is_none() { self.control_variate = Some(vec![0.0_f32; num_params]); } // Get FedAvg aggregation first let aggregated = self.fedavg(updates)?; // Update server control variate let num_clients = updates.len() as f32; if let Some(ref mut cv) = self.control_variate { for (i, delta) in aggregated.iter().enumerate() { if i < cv.len() { // Update control variate based on client updates cv[i] += (delta * self.scaffold_server_lr as f32) / num_clients; } } } // Apply control variate correction let mut corrected = aggregated; if let Some(ref cv) = self.control_variate { for (i, delta) in corrected.iter_mut().enumerate() { if i < cv.len() && i < global_model.weights.len() { *delta += cv[i] * self.scaffold_server_lr as f32; } } } Ok(corrected) } /// FedAdam: Adaptive learning rate aggregation. fn fedadam(&self, updates: &[LocalUpdate]) -> Result, FederatedError> { // Simplified FedAdam: FedAvg with adaptive scaling let mut aggregated = self.fedavg(updates)?; // Apply adaptive scaling based on gradient magnitude let beta1 = 0.9_f32; let beta2 = 0.999_f32; let epsilon = 1e-8_f32; // Compute first and second moment estimates let mut m = vec![0.0_f32; aggregated.len()]; let mut v = vec![0.0_f32; aggregated.len()]; for (i, delta) in aggregated.iter().enumerate() { m[i] = beta1 * m[i] + (1.0 - beta1) * delta; v[i] = beta2 * v[i] + (1.0 - beta2) * delta * delta; } // Apply Adam update for (i, delta) in aggregated.iter_mut().enumerate() { *delta = m[i] / (v[i].sqrt() + epsilon); } Ok(aggregated) } /// FedYogi: Adaptive aggregation with momentum. fn fedyogi(&self, updates: &[LocalUpdate]) -> Result, FederatedError> { // Simplified FedYogi: Similar to FedAdam with different update rule let mut aggregated = self.fedavg(updates)?; let beta1 = 0.9_f32; let beta2 = 0.99_f32; let epsilon = 1e-3_f32; let mut m = vec![0.0_f32; aggregated.len()]; let mut v = vec![0.01_f32; aggregated.len()]; // Initialize with small value for (i, delta) in aggregated.iter().enumerate() { m[i] = beta1 * m[i] + (1.0 - beta1) * delta; // Yogi uses sign of (g^2 - v) for update let g_sq = delta * delta; let sign = if g_sq > v[i] { 1.0 } else { -1.0 }; v[i] += sign * (1.0 - beta2) * g_sq; } for (i, delta) in aggregated.iter_mut().enumerate() { *delta = m[i] / (v[i].sqrt() + epsilon); } Ok(aggregated) } /// Trimmed Mean: Robust aggregation that removes outliers. fn trimmed_mean(&self, updates: &[LocalUpdate]) -> Result, FederatedError> { if updates.is_empty() { return Err(FederatedError::AggregationFailed( "No updates for trimmed mean".to_string(), )); } let num_params = updates[0].weight_deltas.len(); let n = updates.len(); let trim_count = ((n as f32 * self.trim_fraction) as usize).max(1); if n <= 2 * trim_count { // Not enough updates to trim, fall back to FedAvg return self.fedavg(updates); } let mut aggregated = vec![0.0_f32; num_params]; for i in 0..num_params { // Collect all values for this parameter let mut values: Vec = updates .iter() .filter_map(|u| u.weight_deltas.get(i).copied()) .collect(); if values.len() <= 2 * trim_count { // Not enough values, use mean aggregated[i] = values.iter().sum::() / values.len() as f32; } else { // Sort and trim values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); let trimmed = &values[trim_count..values.len() - trim_count]; aggregated[i] = trimmed.iter().sum::() / trimmed.len() as f32; } } Ok(aggregated) } /// Median: Coordinate-wise median for robustness. fn median_aggregate(&self, updates: &[LocalUpdate]) -> Result, FederatedError> { if updates.is_empty() { return Err(FederatedError::AggregationFailed( "No updates for median".to_string(), )); } let num_params = updates[0].weight_deltas.len(); let mut aggregated = vec![0.0_f32; num_params]; for i in 0..num_params { let mut values: Vec = updates .iter() .filter_map(|u| u.weight_deltas.get(i).copied()) .collect(); if values.is_empty() { aggregated[i] = 0.0; } else { values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); let mid = values.len() / 2; aggregated[i] = if values.len().is_multiple_of(2) { f32::midpoint(values[mid - 1], values[mid]) } else { values[mid] }; } } Ok(aggregated) } /// Krum: Byzantine-resilient aggregation. fn krum(&self, updates: &[LocalUpdate]) -> Result, FederatedError> { if updates.is_empty() { return Err(FederatedError::AggregationFailed( "No updates for Krum".to_string(), )); } let n = updates.len(); if n <= 2 { // Fall back to FedAvg for small number of clients return self.fedavg(updates); } // Number of Byzantine clients we want to tolerate let f = (n as f32 * 0.1).ceil() as usize; // Tolerate 10% Byzantine let k = n - f - 2; if k == 0 { return self.fedavg(updates); } // Compute pairwise distances let mut distances: Vec> = vec![vec![0.0; n]; n]; for i in 0..n { for j in (i + 1)..n { let dist = self.euclidean_distance(&updates[i].weight_deltas, &updates[j].weight_deltas); distances[i][j] = dist; distances[j][i] = dist; } } // Compute Krum scores (sum of k closest distances) let mut scores: Vec<(usize, f32)> = Vec::with_capacity(n); for i in 0..n { let mut dists: Vec = distances[i].clone(); dists.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); let score: f32 = dists.iter().take(k + 1).sum(); // +1 because distance to self is 0 scores.push((i, score)); } // Select the update with minimum score scores.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal)); let best_idx = scores[0].0; Ok(updates[best_idx].weight_deltas.clone()) } /// Compute Euclidean distance between two weight vectors. fn euclidean_distance(&self, a: &[f32], b: &[f32]) -> f32 { a.iter() .zip(b.iter()) .map(|(x, y)| (x - y).powi(2)) .sum::() .sqrt() } /// Get the current strategy. #[must_use] pub fn strategy(&self) -> AggregationStrategy { self.strategy } /// Set the trim fraction for trimmed mean. pub fn set_trim_fraction(&mut self, fraction: f32) { self.trim_fraction = fraction.clamp(0.0, 0.4); } /// Set the FedProx mu parameter. pub fn set_fedprox_mu(&mut self, mu: f64) { self.fedprox_mu = mu; } /// Set the Scaffold server learning rate. pub fn set_scaffold_server_lr(&mut self, lr: f64) { self.scaffold_server_lr = lr; } } /// Secure Aggregation using Shamir's Secret Sharing (simplified). #[derive(Debug)] pub struct SecureAggregation { /// Threshold for secret sharing. threshold: usize, /// Number of shares. num_shares: usize, } impl Default for SecureAggregation { fn default() -> Self { Self::new(3, 5) } } impl SecureAggregation { /// Create a new secure aggregation instance. #[must_use] pub fn new(threshold: usize, num_shares: usize) -> Self { Self { threshold, num_shares, } } /// Generate shares for a secret (simplified version). #[must_use] pub fn generate_shares(&self, secret: f32) -> Vec<(usize, f32)> { use rand::Rng; let mut rng = rand::thread_rng(); // Generate random coefficients for polynomial let mut coeffs = vec![secret]; for _ in 1..self.threshold { coeffs.push(rng.gen_range(-1.0_f32..1.0_f32)); } // Evaluate polynomial at different points (1..=self.num_shares) .map(|x| { let y = coeffs .iter() .enumerate() .map(|(i, c)| c * (x as f32).powi(i as i32)) .sum(); (x, y) }) .collect() } /// Reconstruct secret from shares (simplified Lagrange interpolation). #[must_use] pub fn reconstruct(&self, shares: &[(usize, f32)]) -> f32 { if shares.len() < self.threshold { return 0.0; // Not enough shares } let shares = &shares[..self.threshold]; let mut result = 0.0_f32; for (i, &(xi, yi)) in shares.iter().enumerate() { let mut term = yi; for (j, &(xj, _)) in shares.iter().enumerate() { if i != j { let xi_f = xi as f32; let xj_f = xj as f32; term *= -xj_f / (xi_f - xj_f); } } result += term; } result } /// Securely aggregate weight updates. pub fn aggregate(&self, updates: &[LocalUpdate]) -> Result, FederatedError> { if updates.is_empty() { return Err(FederatedError::AggregationFailed( "No updates for secure aggregation".to_string(), )); } let num_params = updates[0].weight_deltas.len(); let mut aggregated = vec![0.0_f32; num_params]; // In practice, each client would generate shares and send to others // Here we simulate the final aggregation step let total_samples: usize = updates.iter().map(|u| u.num_samples).sum(); for i in 0..num_params { for update in updates { if let Some(&delta) = update.weight_deltas.get(i) { let weight = update.num_samples as f32 / total_samples as f32; aggregated[i] += delta * weight; } } } Ok(aggregated) } } #[cfg(test)] mod tests { use super::*; fn sample_updates(num_updates: usize, num_params: usize) -> Vec { (0..num_updates) .map(|i| LocalUpdate { client_id: format!("client_{}", i), round: 1, weight_deltas: vec![0.1 * (i as f32 + 1.0); num_params], num_samples: 100 * (i + 1), local_loss: 0.5, local_accuracy: 0.8, training_time: 10.0, control_variates: None, }) .collect() } #[test] fn test_fedavg() { let aggregator = Aggregator::new(AggregationStrategy::FedAvg); let updates = sample_updates(3, 10); let _global = GlobalModel::default(); let result = aggregator.fedavg(&updates); assert!(result.is_ok()); let aggregated = result.unwrap(); assert_eq!(aggregated.len(), 10); } #[test] fn test_fedprox() { let mut aggregator = Aggregator::fedprox(0.01); let mut updates = sample_updates(3, 10); let global = GlobalModel::default(); let result = aggregator.aggregate(&mut updates, &global); assert!(result.is_ok()); } #[test] fn test_scaffold() { let mut aggregator = Aggregator::scaffold(1.0); let mut updates = sample_updates(3, 10); let global = GlobalModel::default(); let result = aggregator.aggregate(&mut updates, &global); assert!(result.is_ok()); } #[test] fn test_trimmed_mean() { let aggregator = Aggregator::new(AggregationStrategy::TrimmedMean); let updates = sample_updates(5, 10); let result = aggregator.trimmed_mean(&updates); assert!(result.is_ok()); } #[test] fn test_median() { let aggregator = Aggregator::new(AggregationStrategy::Median); let updates = sample_updates(5, 10); let result = aggregator.median_aggregate(&updates); assert!(result.is_ok()); } #[test] fn test_krum() { let aggregator = Aggregator::new(AggregationStrategy::Krum); let updates = sample_updates(5, 10); let result = aggregator.krum(&updates); assert!(result.is_ok()); } #[test] fn test_empty_updates() { let aggregator = Aggregator::new(AggregationStrategy::FedAvg); let updates: Vec = vec![]; let result = aggregator.fedavg(&updates); assert!(result.is_err()); } #[test] fn test_secure_aggregation() { let secure_agg = SecureAggregation::new(3, 5); let updates = sample_updates(3, 10); let result = secure_agg.aggregate(&updates); assert!(result.is_ok()); } #[test] fn test_secret_sharing() { let secure_agg = SecureAggregation::new(3, 5); let secret = 42.0_f32; let shares = secure_agg.generate_shares(secret); assert_eq!(shares.len(), 5); let reconstructed = secure_agg.reconstruct(&shares); assert!((reconstructed - secret).abs() < 0.01); } #[test] fn test_aggregator_strategy() { let aggregator = Aggregator::new(AggregationStrategy::FedProx); assert_eq!(aggregator.strategy(), AggregationStrategy::FedProx); } #[test] fn test_set_parameters() { let mut aggregator = Aggregator::default(); aggregator.set_trim_fraction(0.2); aggregator.set_fedprox_mu(0.05); aggregator.set_scaffold_server_lr(0.5); } #[test] fn test_fedadam() { let aggregator = Aggregator::new(AggregationStrategy::FedAdam); let updates = sample_updates(3, 10); let result = aggregator.fedadam(&updates); assert!(result.is_ok()); } #[test] fn test_fedyogi() { let aggregator = Aggregator::new(AggregationStrategy::FedYogi); let updates = sample_updates(3, 10); let result = aggregator.fedyogi(&updates); assert!(result.is_ok()); } }