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")

Go deeper with CFSGPT

Get AI-powered deep-dives, save terms, and run advanced simulations — free account.

Create Free Account