Federated Learning Distributed Privacy-Preserving Training
# Federated Learning: Distributed Privacy-Preserving Training
## Introduction & Motivation
Federated learning trains on decentralized data; clients compute gradients, server aggregates. Data never centralized; privacy-by-design. Applications: mobile devices, healthcare, finance. Communication-efficient: compress gradients, reduce rounds.
Motivation: Data privacy critical (GDPR, HIPAA). Federated avoids data collection. Handles heterogeneous data distributions across clients.
Applications: On-device ML, healthcare networks, financial institutions, IoT.
---
## Core Concepts & Theory
### Client-Server Architecture
Clients: local data, local training. Server: aggregate gradients, distribute updates.
### Federated Averaging (FedAvg)
Weighted average of client gradient updates; convergence analysis.
### Data Heterogeneity
Non-IID data across clients; statistical challenge. Impacts convergence.
---
## Mathematical Formulation
FedAvg update:
$$\mathbf{w}_{t+1} = \mathbf{w}_t - \eta \sum_{k=1}^K \frac{n_k}{n} \mathbf{g}_k^t$$
where g_k^t = gradient on client k, n_k = client data size.
Convergence with heterogeneity:
$$E[\|
abla F(\mathbf{w}_T)\|^2] = O\left(\frac{1}{T} + \frac{\sigma_h^2}{K}
ight)$$
σ_h² = heterogeneity term; trades off communication vs. accuracy.
---
## Advanced Theory & Extensions
### Personalized Federated Learning
Client-specific models; shared + local parameters.
### Federated Transfer Learning
Pre-train on federated data; transfer downstream.
### Gradient Compression
Quantization, sparsification reduce communication.
---
## Computational Considerations
Communication: Dominant cost; R rounds × K clients × param size.
Local computation: E epochs per client; scales linearly.
Memory: Client must fit model; limits on-device deployment.
---
## Practical Implementation Strategies
### Client Selection
Sample subset K < N clients per round; reduces communication.
### Local Batch Size
Small batches (16-32) on heterogeneous devices.
### Synchronization Timeout
Drop stragglers; avoid waiting for slowest client.
---
## Benchmark Datasets & Evaluation
MNIST Federated: Simulate 1000 clients; IID vs. non-IID settings.
CIFAR-10 Federated: Dirichlet label distribution; tunable α.
Metrics: Convergence rate, communication cost, final accuracy.
---
## Key Challenges & Limitations
### Non-IID Data
Data distributions vary by client; slower convergence.
### Communication Cost
Rounds dominate runtime; gradient compression essential.
### System Heterogeneity
Devices vary in compute/network; stragglers block progress.
---
## Hyperparameter Tuning
Local epochs E: 1-10; more → communication reduction.
Client fraction K/N: 0.01-1.0; typical 0.1.
Learning rate η: 0.01-1.0; adapt by round.
---
## Real-World Applications & Case Studies
Keyboard Prediction (Google): On-device federated learning; privacy-preserving.
Medical Imaging: Multi-hospital federated training; no data sharing.
Mobile Analytics: Device-level model updates; privacy-first.
---
## Integration with Other Methods
Federated + Differential Privacy → add DP noise to gradients.
Federated + Compression → quantize/sparsify before communication.
---
## Summary & Key Takeaways
Federated learning distributes training across clients, maintaining privacy while aggregating model updates via gradient averaging. Handles heterogeneous data and communication constraints.
Principles:
1. FedAvg: weighted gradient averaging across clients.
2. Communication-efficient: compress gradients, async updates.
3. Non-IID data challenges convergence; requires regularization.
4. Client selection, timeout strategies manage system heterogeneity.
5. Privacy-by-design; no central data repository.
---
---
## Appendix: Practical Labs
### Lab 1: Gradient Aggregation
import torch
import numpy as np
def federated_average(client_gradients, client_data_sizes):
"""FedAvg: weighted average of client gradients"""
total_samples = sum(client_data_sizes)
weights = [n / total_samples for n in client_data_sizes]
aggregated_grad = None
for w, grad in zip(weights, client_gradients):
if aggregated_grad is None:
aggregated_grad = w * grad
else:
aggregated_grad += w * grad
return aggregated_grad
# Test
grads = [torch.randn(10, 5) for _ in range(4)]
sizes = [100, 150, 120, 80]
agg_grad = federated_average(grads, sizes)
print(f"Aggregated gradient shape: {agg_grad.shape}")
assert agg_grad.shape == (10, 5), "Should preserve shape"
assert not torch.allclose(agg_grad, torch.zeros_like(agg_grad)), "Should be non-trivial"
print("✓ Gradient aggregation working")
if __name__ == "__main__":
print("Lab 1: Aggregation - PASSED")### Lab 2: Local SGD Step
import torch
import torch.nn as nn
import torch.optim as optim
def client_local_update(model, x, y, optimizer, n_epochs=1):
"""Local SGD on client device"""
model.train()
criterion = nn.CrossEntropyLoss()
losses = []
for epoch in range(n_epochs):
optimizer.zero_grad()
logits = model(x)
loss = criterion(logits, y)
loss.backward()
optimizer.step()
losses.append(loss.item())
return losses[-1]
# Test
model = nn.Linear(10, 5)
optimizer = optim.SGD(model.parameters(), lr=0.01)
x = torch.randn(32, 10)
y = torch.randint(0, 5, (32,))
loss = client_local_update(model, x, y, optimizer, n_epochs=2)
print(f"Client loss: {loss:.4f}")
assert loss > 0, "Loss should be positive"
print("✓ Local update working")
if __name__ == "__main__":
print("Lab 2: Local Update - PASSED")### Lab 3: Non-IID Data Simulation
import torch
import numpy as np
def create_noniid_data(n_classes=10, n_clients=5, samples_per_client=100, alpha=0.1):
"""Create non-IID data: Dirichlet distribution over classes"""
# Dirichlet: lower alpha → more heterogeneous
label_distribution = np.random.dirichlet([alpha] * n_classes, size=n_clients)
client_data = []
for c in range(n_clients):
# Sample classes according to distribution
class_counts = (label_distribution[c] * samples_per_client).astype(int)
labels = []
for cls, count in enumerate(class_counts):
labels.extend([cls] * count)
client_data.append({
'labels': torch.tensor(labels[:samples_per_client]),
'distribution': label_distribution[c]
})
return client_data
# Test
clients = create_noniid_data(n_classes=10, n_clients=5, alpha=0.5)
print(f"Client 0 class distribution: {clients[0]['distribution'][:3]}")
assert len(clients) == 5, "Should have 5 clients"
assert len(clients[0]['labels']) == 100, "Should have samples per client"
print("✓ Non-IID data simulation working")
if __name__ == "__main__":
print("Lab 3: Non-IID - PASSED")### Lab 4: Communication Cost
import torch
import numpy as np
def compute_communication_cost(model_params, n_rounds, n_clients, compression_ratio=1.0):
"""Estimate communication cost in gradient exchanges"""
total_params = sum(p.numel() for p in model_params)
# Bits per param (32-bit float)
bytes_per_param = 4
# Communication per round: clients send gradients to server + server broadcasts
bytes_per_round = total_params * bytes_per_param * n_clients * 2 / compression_ratio
total_bytes = bytes_per_round * n_rounds
return {
'total_params': total_params,
'bytes_per_round': bytes_per_round,
'total_mb': total_bytes / 1e6
}
# Test
model = torch.nn.Linear(100, 50)
cost = compute_communication_cost(model.parameters(), n_rounds=100, n_clients=1000, compression_ratio=1.0)
print(f"Total params: {cost['total_params']}, Communication: {cost['total_mb']:.1f} MB")
assert cost['total_params'] > 0, "Should have parameters"
assert cost['total_mb'] > 0, "Should have communication cost"
print("✓ Communication cost estimation working")
if __name__ == "__main__":
print("Lab 4: Communication - PASSED")