Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
+636
View File
@@ -0,0 +1,636 @@
//! 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<Vec<f32>>,
/// 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<Vec<f32>, 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<Vec<f32>, 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<Vec<f32>, 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<Vec<f32>, 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<Vec<f32>, 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<Vec<f32>, 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<Vec<f32>, 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<f32> = 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::<f32>() / 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::<f32>() / trimmed.len() as f32;
}
}
Ok(aggregated)
}
/// Median: Coordinate-wise median for robustness.
fn median_aggregate(&self, updates: &[LocalUpdate]) -> Result<Vec<f32>, 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<f32> = 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<Vec<f32>, 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<f32>> = 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<f32> = 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::<f32>()
.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<Vec<f32>, 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<LocalUpdate> {
(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<LocalUpdate> = 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());
}
}