Initial commit
This commit is contained in:
@@ -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());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user