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

---

Go deeper with CFSGPT

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

Create Free Account