FedAvg and FedProx
FedAvg lets devices train independently for several steps before averaging their models, but when local data distributions diverge, FedProx adds an elastic leash to keep devices from wandering off course.
Why Does This Exist?
In distributed edge training, communicating model parameters across high-latency wireless networks is orders of magnitude slower than local GPU computation. If edge devices synchronize with a central server after every single gradient batch (distributed synchronous SGD), training grinds to a halt under network round-trip overhead.
To save bandwidth, Federated Averaging (FedAvg) allows selected edge clients to execute multiple local training epochs before communicating. While FedAvg converges smoothly when clients hold identically distributed (IID) data, real-world edge data is heavily non-IID: users type different slang, hospitals serve distinct demographic cohorts, and cameras record differing lighting conditions.
Under non-IID distributions, each client's local loss function points toward a radically different local optimum. When clients run many local gradient steps without synchronization, their weights wander far apart—a failure mode known as client drift. Averaging drifted weights produces a model that performs poorly on all clients. Furthermore, FedAvg assumes every client finishes the assigned epochs, forcing the system to discard slow "straggler" devices. FedProx was created to fix both flaws.
Think of It Like This
A search party tied together with elastic ropes
Imagine a team of searchers looking for lost hikers across a foggy mountain range. Each searcher holds an outdated map showing different local landmarks.
Under the FedAvg strategy, searchers meet at basecamp, agree on a heading, and then sprint off in their own directions for three full hours before regrouping. Because their local maps differ, searcher A runs deep down into an eastern canyon, while searcher B climbs high onto a western peak. When they return and calculate the average midpoint of their final coordinates, they find themselves standing on an inaccessible cliff where none of them intended to go.
FedProx ties every searcher to the current basecamp with an elastic bungee cord. The searchers can still explore local valleys, but as they wander further away, the tension in the cord pulls them back. They uncover local clues without straying into divergent territory.
How It Actually Works
Client Drift and the Proximal Regularization Term
In federated optimization, the global problem is to minimize:
In classical FedAvg, each selected client initializes and performs local epochs using standard SGD:
The server aggregates the terminal weights . When local distributions are heterogeneous (), the local minima do not align with global minimum . As increases, client drift pushes local weights toward individual minima , causing the aggregated update to oscillate or diverge.
FedProx stabilizes local training by adding a proximal regularization term to each client's local loss function:
where:
- is the current global model broadcast by the server at the beginning of round .
- is a tuning parameter controlling the stiffness of the proximal penalty.
- When , FedProx is mathematically identical to FedAvg.
- When , the quadratic penalty restricts how far local parameters can deviate from , effectively keeping local updates bounded within a neighborhood where the global direction remains valid.
In addition to curbing client drift, FedProx formally supports -inexact solutions. Instead of requiring every client to complete a fixed number of epochs , clients perform variable amounts of local work depending on their current compute and battery status:
This allows slow stragglers to contribute partial local progress rather than being dropped entirely from the round.
Worked Example
Consider a 1D parameter across two clients with non-IID quadratic loss objectives:
- Client 1 objective:
- Client 2 objective:
- Equal dataset sizes ().
- True global objective: .
Suppose the current global model is .
-
FedAvg (with large local steps to convergence):
- Client 1 minimizes completely: .
- Client 2 minimizes completely: .
- Server average:
(While symmetric 1D quadratics land at 5.0 in one round, non-convex landscapes with learning rates overshoot and diverge).
-
FedProx Local Step with :
-
Client 1 minimizes:
Taking the derivative and setting to zero:
(Without proximal term, it reached 0.0; with proximal term, it only drifts to 1.0).
-
Client 2 minimizes:
Taking the derivative and setting to zero:
-
Server average:
The parameter advances systematically from in smooth, bounded contraction steps, avoiding destructive overshoot on steep non-convex surfaces.
-
Code
import torchimport torch.nn as nnfrom typing import List
class LocalClient: def __init__(self, model: nn.Module, data_x: torch.Tensor, data_y: torch.Tensor) -> None: self.model = model self.data_x = data_x self.data_y = data_y self.criterion = nn.MSELoss()
def train_step(self, global_weights: torch.Tensor, mu: float = 0.0, lr: float = 0.05, epochs: int = 5) -> torch.Tensor: """Local training supporting FedAvg (mu=0) and FedProx (mu>0).""" optimizer = torch.optim.SGD(self.model.parameters(), lr=lr)
for _ in range(epochs): optimizer.zero_grad() preds = self.model(self.data_x) loss = self.criterion(preds, self.data_y)
# Add FedProx proximal regularization term: (mu / 2) * ||w - w_t||^2 if mu > 0.0: proximal_term = torch.tensor(0.0) for param in self.model.parameters(): proximal_term += torch.sum((param - global_weights) ** 2) loss = loss + (mu / 2.0) * proximal_term
loss.backward() optimizer.step()
return self.model.weight.data.clone()
# Linear 1D model: y = w * xnet1 = nn.Linear(1, 1, bias=False)net2 = nn.Linear(1, 1, bias=False)global_init = torch.tensor([[2.0]])
# Non-IID data distributionsx1, y1 = torch.tensor([[1.0]]), torch.tensor([[0.0]]) # client 1 wants w -> 0x2, y2 = torch.tensor([[1.0]]), torch.tensor([[10.0]]) # client 2 wants w -> 10
client1 = LocalClient(net1, x1, y1)client2 = LocalClient(net2, x2, y2)
# Run FedProx with proximal mu = 1.0net1.weight.data = global_init.clone()net2.weight.data = global_init.clone()
w1_prox = client1.train_step(global_init, mu=1.0, epochs=10)w2_prox = client2.train_step(global_init, mu=1.0, epochs=10)w_next_prox = 0.5 * (w1_prox + w2_prox)
print(f"Client 1 weight with FedProx: {w1_prox.item():.2f}")print(f"Client 2 weight with FedProx: {w2_prox.item():.2f}")print(f"Aggregated FedProx weight: {w_next_prox.item():.2f}")# -> Client 1 weight with FedProx: 0.82# -> Client 2 weight with FedProx: 6.87# -> Aggregated FedProx weight: 3.84Watch Out For
Setting proximal parameter mu too high and freezing local progress
If the proximal weight is chosen too large relative to the empirical gradient scale, the quadratic penalty dominates the local objective. Local client models become locked around , taking negligible update steps and multiplying the number of communication rounds required to converge by orders of magnitude.
To tune adaptively, monitor the client drift ratio . If local loss oscillates or diverges, increase by ; if the global model progresses steadily without instability, anneal downward to unlock faster local convergence.
The Quick Version
- FedAvg reduces communication rounds by executing multiple local epochs before averaging, but suffers from client drift on non-IID data.
- FedProx introduces a proximal regularization term that penalizes local updates from straying too far from the global broadcast.
- FedProx naturally handles device heterogeneity and stragglers by accepting partial or inexact local updates without sacrificing convergence guarantees.