Hospitals want to build better diagnostic models, but cannot share patient records. Phones could improve keyboard predictions, but users' messages should never leave their devices. Organisations in different countries hold data under different legal regimes. Federated learning (FL) offers a way to train a shared model without collecting raw data in one place: the model travels to the data, not the other way round.
The setting#
- A server coordinates training of a global model.
- Many clients (phones, hospitals, organisations) each hold private local data.
- Clients train locally and send model updates (weights or gradients), not data, back to the server.
Cross-device FL: millions of phones, each with little data, intermittently available (e.g. Google's Gboard next-word prediction, one of the first large deployments). Cross-silo FL: a handful of institutions (hospitals, banks) with large datasets and reliable connections.
FedAvg#
McMahan et al. (2017) proposed Federated Averaging, the standard algorithm:
- The server sends the current global weights $\mathbf{w}^t$ to a sample of clients.
- Each selected client $k$ runs several epochs of SGD on its local data, producing $\mathbf{w}_k^{t+1}$.
- The server averages the updates, weighted by client dataset size $n_k$:
- Repeat for many rounds.
Doing multiple local steps before communicating dramatically reduces communication compared with sending every gradient.
import copy
import numpy as np
import torch, torch.nn as nn, torch.nn.functional as F
def local_train(model, X, y, epochs=2, lr=0.1):
model = copy.deepcopy(model); opt = torch.optim.SGD(model.parameters(), lr=lr)
for _ in range(epochs):
for i in range(0, len(X), 32):
loss = F.cross_entropy(model(X[i:i + 32]), y[i:i + 32])
opt.zero_grad(); loss.backward(); opt.step()
return model.state_dict(), len(X)
def fedavg(states_sizes):
total = sum(n for _, n in states_sizes)
return {k: sum(s[k] * (n / total) for s, n in states_sizes) for k in states_sizes[0][0]}
torch.manual_seed(0); rng = np.random.default_rng(0)
global_model = nn.Sequential(nn.Linear(10, 32), nn.ReLU(), nn.Linear(32, 2))
# Five clients with non-IID data: each has a different class balance
clients = []
for k in range(5):
X = torch.randn(400, 10); w_true = torch.ones(10)
y = ((X @ w_true + rng.normal(0, 1, 400).astype(np.float32) + (k - 2)) > 0).long()
clients.append((X, y))
for rnd in range(20):
chosen = rng.choice(5, 3, replace=False) # partial participation
updates = [local_train(global_model, *clients[k]) for k in chosen]
global_model.load_state_dict(fedavg(updates))
Xt = torch.cat([c[0] for c in clients]); yt = torch.cat([c[1] for c in clients])
print("global accuracy:", (global_model(Xt).argmax(1) == yt).float().mean().item())Challenges#
- Non-IID data: clients' data distributions differ (different languages, regions, patient populations). Local models drift apart ("client drift"), slowing or degrading convergence. Remedies: FedProx (a proximal term keeping local models near the global model), SCAFFOLD (control variates), and personalisation (fine-tuning the global model per client, or mixing global and local components).
- Communication cost: models can be large; networks slow. Use fewer rounds with more local work, update compression (quantisation, sparsification), and smaller models.
- Systems heterogeneity: devices differ in compute, battery and connectivity; clients drop out.
- Fairness: the global model may serve large or typical clients well and small or unusual ones poorly; fairness-aware aggregation can help.
- Robustness: malicious clients can send poisoned updates (backdoors); robust aggregation (median, trimmed mean, Krum) mitigates some attacks.
Privacy: FL is not automatically private#
Applications#
- Mobile: keyboard prediction, voice-assistant improvements, on-device personalisation.
- Healthcare: multi-hospital models for medical imaging and outcome prediction; studies such as a multi-institution COVID-19 outcome model (EXAM, 2021) trained across many hospitals worldwide.
- Finance: fraud detection across institutions.
- Cross-organisational and cross-border collaboration, where legal constraints prevent data pooling.
Frameworks: Flower, TensorFlow Federated, NVIDIA FLARE, PySyft, FedML.