Federated and Decentralized Learning
Instead of collecting everyone's private data onto one central supercomputer, send the model to edge devices, train locally, and only exchange mathematical updates.
Why Does This Exist?
Traditional deep learning pools all training data onto a central cluster or cloud datacenter. In modern domains—such as next-word prediction on mobile smartphones, medical diagnosis across competing hospitals, or financial fraud detection across commercial banks—centralizing raw data is legally prohibited by privacy statutes (such as GDPR, CCPA, and HIPAA) or commercially unacceptable due to trade secrets.
Moreover, streaming gigabytes of raw sensory data from millions of edge devices creates unsustainable network bandwidth bottlenecks and consumes excessive device battery.
Federated Learning (FL) reverses the computation paradigm: it brings the code to the data rather than the data to the code. Edge devices compute local gradient steps on their private datasets and upload only lightweight parameter deltas to an aggregation coordinator. Fully Decentralized Learning takes this a step further by removing the central coordinator entirely, diffusing knowledge across peer-to-peer mesh networks via gossip consensus algorithms.
Think of It Like This
A league of master chefs refining a secret sauce
Imagine ten independent restaurant chefs across different countries trying to perfect a universal pasta sauce recipe. None of the chefs are willing to share their private supplier lists, proprietary spice quantities, or confidential family heritage with a central food corporation.
Instead, an organizer sends the baseline recipe card to all ten kitchens. Each chef cooks the recipe locally, tests it with their local customers, and writes down small recommended adjustments: "add two pinches of basil, reduce salt by one gram."
The chefs mail only their adjustment cards to an impartial clerk who averages the recommendations and publishes Revision 2. In a fully decentralized kitchen league, there is no clerk at all: chefs simply phone their two nearest neighboring restaurants after every evening service and average their notes directly.
How It Actually Works
Hub-and-Spoke Aggregation vs Peer-to-Peer Consensus
In a standard centralized Federated Learning (Hub-and-Spoke) system, clients participate under the direction of a central aggregation server. Client possesses private dataset containing samples, with total dataset size .
The objective is to minimize the global empirical risk:
The lifecycle of each federated round proceeds through four phases:
- Selection: The server samples a random fraction of clients ().
- Broadcast: The server transmits current global parameters to each chosen client.
- Local Training: Each client initializes local model and runs epochs of SGD over local minibatches of size : yielding total local update .
- Aggregation: Selected clients transmit updates back to the server (often protected by Secure Aggregation protocols or Differential Privacy noise). The server computes the weighted average:
In Fully Decentralized Learning (Peer-to-Peer), there is no central server. Clients are nodes in a communication graph . Each node maintains its own local parameters and communicates solely with its direct neighbors .
Decentralized consensus is governed by a doubly-stochastic mixing matrix satisfying , , and if . At each iteration :
Decentralized algorithms (such as D-PSGD and EXTRA) eliminate single-point server failures, avoid high-concurrency bandwidth bottlenecks at central datacenters, and scale robustly across dynamic peer topologies.
Worked Example
Trace one round of Federated Averaging across 3 hospital nodes with unequal patient counts:
- Total parameters: 2 weights .
- Global initialization: .
- Node sample counts: Hospital 1 has , Hospital 2 has , Hospital 3 has . Total .
- Client weights: .
Each hospital runs local SGD on its private records, arriving at:
- Hospital 1:
- Hospital 2:
- Hospital 3:
The server aggregates the updates using sample weighting:
The updated global model for round 1 is:
Hospital 3 had the largest dataset (), so its local trajectory heavily steers the global consensus.
Code
import torchimport torch.nn as nnfrom typing import List, Dict
class SimpleNet(nn.Module): def __init__(self) -> None: super().__init__() self.fc = nn.Linear(2, 1, bias=False) self.fc.weight.data = torch.tensor([[10.0, 5.0]])
def forward(self, x: torch.Tensor) -> torch.Tensor: return self.fc(x)
def federated_averaging( global_model: nn.Module, client_weights: List[Dict[str, torch.Tensor]], client_sample_sizes: List[int]) -> None: """Aggregate client weights weighted by sample sizes.""" total_samples = sum(client_sample_sizes) fractions = [n / total_samples for n in client_sample_sizes]
# Initialize zeroed accumulator aggregated_dict = { key: torch.zeros_like(param) for key, param in global_model.state_dict().items() }
# Weighted sum of client state dicts for client_state, frac in zip(client_weights, fractions): for key in aggregated_dict: aggregated_dict[key] += frac * client_state[key]
# Load aggregated weights into global model global_model.load_state_dict(aggregated_dict)
# Setup 3 client models with local resultsclient_states = [ {"fc.weight": torch.tensor([[10.5, 4.0]])}, {"fc.weight": torch.tensor([[9.0, 6.0]])}, {"fc.weight": torch.tensor([[11.0, 5.5]])}]sample_counts = [100, 300, 600]
server_model = SimpleNet()federated_averaging(server_model, client_states, sample_counts)
print("Aggregated Global Weights:")print(server_model.fc.weight.data)# -> Aggregated Global Weights:# -> tensor([[10.3500, 5.5000]])Watch Out For
Data leakage via gradient inversion attacks
While federated learning avoids transmitting raw data, transmitting raw gradients is not inherently private. Malicious servers or sniffing adversaries can run gradient inversion (such as Deep Leakage from Gradients / DLG), optimizing synthetic inputs until their calculated gradients match the client's uploaded vector, reconstructing private images and text verbatim.
To prevent reconstruction attacks, always combine federated aggregation with Differential Privacy (DP-SGD)—clipping gradient norms and adding calibrated Gaussian noise—or utilize cryptographic Secure Multi-Party Computation (SMPC) where the server can only decrypt the sum of client updates, never any individual client's gradient.
The Quick Version
- Keeps raw training data decentralized on edge devices, transmitting only model parameter updates or gradients.
- Centralized federated learning uses a hub-and-spoke coordinator with sample-weighted aggregation rounds.
- Decentralized peer-to-peer learning eliminates server bottlenecks by exchanging model states across graph neighbors using doubly-stochastic mixing matrices.