FedProx 与 SCAFFOLD 算法
FedAvg 在 Non-IID 数据下收敛性下降。FedProx 和 SCAFFOLD 是两种重要的改进算法,分别通过近端正则化和方差缩减来处理数据异构性。
FedProx
FedProx 在本地目标函数中添加近端正则项,限制本地模型不偏离全局模型太远:
min_w F_k(w) + (μ/2) × ||w - w_t||²
其中:
- F_k(w): 客户端 k 的本地损失
- w_t: 当前全局模型
- μ: 近端项系数(控制偏离程度)python
class FedProxClient:
def __init__(self, model, mu=0.01):
self.model = model
self.mu = mu
def local_train(self, global_model, dataloader, lr=0.01, epochs=5):
"""带近端正则项的本地训练"""
for epoch in range(epochs):
for batch in dataloader:
loss = self.compute_loss(batch)
# 标准梯度
grads = torch.autograd.grad(loss, self.model.parameters())
# 近端正则项梯度: μ × (w - w_t)
for param, global_param, grad in zip(
self.model.parameters(),
global_model.parameters(),
grads
):
prox_grad = grad + self.mu * (param - global_param)
param.data -= lr * prox_grad
return self.modelμ 值选择
| μ 值 | 效果 | 适用场景 |
|---|---|---|
| 0 | 退化为 FedAvg | IID 数据 |
| 0.001 | 轻度约束 | 轻度异构 |
| 0.01 | 中等约束 | 中度异构 |
| 0.1 | 强约束 | 高度异构 |
SCAFFOLD
SCAFFOLD 使用控制变量(control variate)来纠正客户端漂移:
python
class SCAFFOLDClient:
def __init__(self, model):
self.model = model
self.c_local = None # 客户端控制变量
def local_train(self, global_model, c_global, dataloader, lr=0.01):
"""SCAFFOLD 本地训练"""
if self.c_local is None:
self.c_local = deepcopy(c_global)
for batch in dataloader:
loss = self.compute_loss(batch)
grads = torch.autograd.grad(loss, self.model.parameters())
# 修正梯度:减去漂移
for param, grad, c_l, c_g in zip(
self.model.parameters(), grads,
self.c_local, c_global
):
corrected_grad = grad - c_l + c_g
param.data -= lr * corrected_grad
# 更新客户端控制变量
self.update_control_variate(c_global)
return self.model对比分析
| 特性 | FedAvg | FedProx | SCAFFOLD |
|---|---|---|---|
| 额外超参数 | 无 | μ | 无 |
| 通信开销 | 低 | 低 | 中(传输控制变量) |
| Non-IID 收敛 | 差 | 好 | 最好 |
| IID 收敛 | 基准 | 略慢于基准 | 等于基准 |
| 实现复杂度 | 低 | 低 | 中 |
算法选择
- IID 或轻度异构:FedAvg 足够
- 中度异构:FedProx(实现简单)
- 高度异构:SCAFFOLD(收敛最好)
- 通信受限:FedProx(无额外通信)
SCAFFOLD 的额外开销
SCAFFOLD 需要传输控制变量(与模型参数同大小),通信量翻倍。在带宽受限场景下需权衡。