⚖️ AI Ethics, Society & Careers · Lecture 6 of 17

Federated Learning: Training Without Centralising Data

Federated learning trains a shared model across many devices or institutions while raw data stays local. We derive FedAvg, discuss non-IID data, communication costs, secure aggregation, privacy limits and real applications.

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:

  1. The server sends the current global weights $\mathbf{w}^t$ to a sample of clients.
  2. Each selected client $k$ runs several epochs of SGD on its local data, producing $\mathbf{w}_k^{t+1}$.
  3. The server averages the updates, weighted by client dataset size $n_k$:
$$ \mathbf{w}^{t+1} = \sum_{k \in S_t}\frac{n_k}{\sum_{j \in S_t}n_j}\,\mathbf{w}_k^{t+1} $$
  1. Repeat for many rounds.

Doing multiple local steps before communicating dramatically reduces communication compared with sending every gradient.

python
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#

  1. 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).
  2. Communication cost: models can be large; networks slow. Use fewer rounds with more local work, update compression (quantisation, sparsification), and smaller models.
  3. Systems heterogeneity: devices differ in compute, battery and connectivity; clients drop out.
  4. Fairness: the global model may serve large or typical clients well and small or unusual ones poorly; fairness-aware aggregation can help.
  5. 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.

JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST · Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

⚖️ AI Ethics, Society & Careers

Privacy in Machine Learning and Differential Privacy

Models can leak the data they were trained on. We examine re-identification and model attacks, why anonymisation often fails, the mathematics of differential privacy, DP-SGD, and practical privacy-by-design.

Advanced⏱ 6 min#261
⚖️ AI Ethics, Society & Careers

AI Safety and Alignment: Making Capable Systems Do What We Intend

As AI systems grow more capable, ensuring they pursue intended goals becomes critical. We cover specification gaming, reward hacking, goal misgeneralisation, current alignment techniques, interpretability, evaluations and governance of frontier models.

Intermediate⏱ 6 min#263
⚖️ AI Ethics, Society & Careers

Explainable AI: LIME, SHAP and Interpretable Models

Why did the model decide that? We distinguish interpretable models from post-hoc explanations, derive Shapley values and SHAP, explain LIME, cover global vs local explanations and counterfactuals, and discuss the limits of explanations.

Intermediate⏱ 5 min#260