Federated Learning - Distributed Training
# Federated Learning - Distributed Training
## Introduction & Motivation
Federated learning: train on decentralized data. Keep data local, share only gradients. Applications: privacy-preserving training, edge learning.
Motivation: Enable training without centralizing data.
Applications: Mobile learning, healthcare, finance.
---
## Core Concepts & Theory
### Local Training
Train on device.
### Gradient Sharing
Send updates to server.
### Aggregation
Average across clients.
### Communication Efficiency
Minimize bandwidth.
---
## Mathematical Formulation
Federated Averaging:
$$ heta_{t+1} = \sum_i p_i heta_i^t$$
Client Update:
$$ heta_i^t = heta^t - \eta
abla L_i( heta^t)$$
Cost:
$$ ext{communication} \propto M \cdot n_{clients}$$
---
## Advanced Theory & Extensions
### Gradient Compression
Reduce communication.
### Personalization
Client-specific models.
### Asynchronous Aggregation
Non-synchronized updates.
---
## Computational Considerations
Client computation: O(n·D²).
Communication: O(D).
Total rounds: Depends on convergence.
---
## Practical Implementation Strategies
### Client Selection
Choose subset per round.
### Compression Strategy
Quantize gradients.
### Variance Reduction
Control client drift.
---
## Benchmark Datasets & Evaluation
MNIST: Federated setting.
CIFAR-10: Non-IID data.
SHAKESPEARE: Text generation.
---
## Key Challenges & Limitations
### Non-IID Data
Heterogeneous distributions.
### Communication Cost
Bandwidth bottleneck.
### Convergence
Slower than centralized.
---
## Hyperparameter Tuning
Learning rate: 0.01-0.1.
Local epochs: 1-10.
Client participation: 10-100%.
---
## Real-World Applications & Case Studies
Mobile Devices: On-device learning.
Healthcare: Preserve patient data.
Finance: Regulatory compliance.
---
## Integration with Other Methods
Federated + differential privacy; + compression.
---
## Summary & Key Takeaways
Federated learning enables distributed training.
Principles:
1. Decentralization: Keep data local.
2. Gradient sharing: Send updates only.
3. Aggregation: Average updates.
4. Communication: Minimize bandwidth.
5. Privacy: Protect individual data.
---
## Appendix: Practical Labs
### Lab 1: Federated Averaging
import numpy as np
def federated_averaging(client_weights, client_participation):
"""Aggregate weights from clients"""
num_clients = len(client_weights)
averaged = np.zeros_like(client_weights[0])
weight_sum = np.sum(client_participation)
for i, weights in enumerate(client_weights):
participation = client_participation[i] / weight_sum
averaged += participation * weights
return averaged
np.random.seed(42)
weights = [np.random.randn(100) for _ in range(10)]
participation = np.ones(10)
avg = federated_averaging(weights, participation)
assert avg.shape == weights[0].shape
print("✓ Federated averaging working")### Lab 2: Gradient Compression
import numpy as np
def compress_gradient(gradient, compression_ratio=0.1):
"""Compress gradient for communication"""
num_keep = max(1, int(gradient.size * compression_ratio))
# Keep top-k magnitudes
flat = gradient.flatten()
top_indices = np.argsort(np.abs(flat))[-num_keep:]
compressed = np.zeros_like(flat)
compressed[top_indices] = flat[top_indices]
return compressed.reshape(gradient.shape)
np.random.seed(42)
grad = np.random.randn(1000)
compressed = compress_gradient(grad, 0.1)
assert np.count_nonzero(compressed) <= 100
print("✓ Gradient compression working")### Lab 3: Client Drift
import numpy as np
def compute_client_drift(local_params, global_params):
"""Measure divergence between local and global"""
drift = np.mean((local_params - global_params) ** 2)
return drift
np.random.seed(42)
local = np.random.randn(100)
global_p = np.random.randn(100)
drift = compute_client_drift(local, global_p)
assert drift >= 0
print(f"✓ Client drift: {drift:.4f}")### Lab 4: Communication Cost
def estimate_communication_cost(num_clients, param_size, num_rounds):
"""Estimate communication cost"""
messages_per_round = num_clients * 2 # upload + download
total_bytes = messages_per_round * param_size * num_rounds
total_mb = total_bytes / (1024 * 1024)
return total_mb
cost = estimate_communication_cost(100, 10000, 100)
assert cost > 0
print(f"✓ Communication cost: {cost:.1f} MB")---