做联邦学习(Federated Learning)这么多年,我越来越确信一件事:联邦聚合(Federated Aggregation)才是整个联邦系统的“魂”。很多人一上来就研究 FedAvg、FedProx、SCAFFOLD 这些算法名字,但真正在线上项目里跑过之后,你会发现聚合算法选错,后面所有隐私保护、加密方案全都白搭。这篇博文我想以实操视角,把三种主流联邦聚合算法的原理、实现、对比和踩坑经验完整讲透,给正要入手联邦学习的工程师一份可以直接落地的参考。
1. 联邦聚合到底在解决什么问题
1.1 没有聚合,联邦学习只是一堆本地模型“各说各话”
联邦学习的核心流程你可能已经听过很多次:中心服务器下发全局模型,参与客户端在本地数据上做几轮训练,然后把模型更新传回服务器,服务器用某种方式把这些更新合并成新的全局模型。这里“合并”这一步,就是联邦聚合。
为什么聚合这么关键?因为参与训练的客户端数据分布往往极其不一致。有的客户端样本量相差几十倍,有的客户端大量样本来自同一个类别,有的设备网络不稳定,每轮参与名单还一直在变。聚合算法需要在这种复杂条件下,把来自不同“局部世界”的模型更新,压缩成一个能代表全局分布的模型更新。如果聚合策略设计得不好,本地模型学得越努力,全局模型反而越糟糕。
你可以把每个客户端想象成一个只在自己教室里参加考试的学生。学生A只做过数学卷子,学生B只做过语文卷子。如果老师只是简单把他们试卷上的错误率平均一下就出最终成绩,那这个成绩既不能代表数学水平,也不能代表语文水平。联邦聚合算法的演进,本质上就是在设计“如何更公平、更正确地合并这些学生各自的经验”。
1.2 三种算法的本质差异,一句话版本
先给一个全貌,后面再逐个拆解:
- FedAvg:最简单,把所有客户端的模型参数按样本比例做加权平均。
- FedProx:在 FedAvg 基础上,给每个客户端本地目标加一个“别离全局模型太远”的近端约束。
- SCAFFOLD:引入控制变量,修正本地更新方向,降低客户端漂移。
一句话总结就是:FedAvg 是基线方案,适用数据分布相对均匀的场景;FedProx 是给“容易跑偏”的本地训练加刹车;SCAFFOLD 是给“方向本身就不对”的本地梯度做纠偏。三者不是替代关系,而是不同复杂度、不同适用场景下的解决方案。
你可能会想,既然 FedAvg 这么简单,为什么还要搞后面两个?因为真实世界的数据分布从来不会配合你的理想,Non-IID 就是联邦学习逃不掉的常态。
2. FedAvg:简单,但别小看它
2.1 算法流程与加权平均公式
FedAvg 最早由 McMahan 等人在 2017 年提出,全称是 Federated Averaging。它的做法直白到让人怀疑人生:
- 服务器初始化一个全局模型 w0。
- 每一轮 t,服务器从所有客户端中随机挑选一部分(比如 10%)参与本轮训练。
- 被选中的客户端下载全局模型 wt,在本地用自己的数据跑若干轮 SGD(一般设为 1 到 10 个 epoch)。
- 每个客户端把本地训练后的模型参数(或模型参数变化量 Δw)传回服务器。
- 服务器把所有客户端传来的模型更新按样本数量做加权平均,得到新的全局模型。
用公式写就是:
w_{t+1} = Σ_{k=1}^{K} (n_k / n) * w_{t+1}^{(k)}
其中 K 是参与客户端数量,n_k 是第 k 个客户端的本地样本量,n = Σ n_k,w_{t+1}^{(k)} 是第 k 个客户端本地训练后的模型参数。
举个具体例子:客户 A 有 100 条样本,训练后模型参数从 [1.0, 0.5] 更新到 [1.2, 0.6];客户 B 有 300 条样本,更新到 [0.9, 0.4]。按样本比例计算,客户 A 权重是 100/400=0.25,客户 B 权重是 300/400=0.75,聚合后全局模型参数就是 0.25*[1.2, 0.6] + 0.75*[0.9, 0.4] = [0.975, 0.45]。
这个加权平均的逻辑很简单:谁的样本多,谁的更新就在全局模型里占更大的话语权。FedAvg 的通信效率优势来自“本地多轮更新”的设计——客户端在本地迭代好几个 epoch 后才传回模型,而不是每步都传梯度。在带宽受限的真实环境里,这一设计能把通信量降低几个数量级,这也是 FedAvg 能成为联邦学习默认 basline 的根本原因。
2.2 为什么 Non-IID 数据下会出现“客户端漂移”
FedAvg 最受诟病的问题,是在 Non-IID(非独立同分布)数据下收敛慢、最终模型精度差。原因在于“本地多轮更新”与“全局平均”之间存在根本性张力。
如果数据是 IID 的,也就是每个客户端的本地数据都像切蛋糕一样均匀地代表了全局分布,那么本地更新就是全局梯度的无偏估计,平均一下非常合理。但 Non-IID 场景下,客户端的本地目标函数和全局目标函数的梯度方向可能差得很远。比如某家医院只收录肺部感染患者影像,另一家只收录骨折患者影像,它们各自在本地学到的方向,对全局模型来说可能是互相“打架”的。
当每个客户端在本地跑多轮 SGD 后,它的模型参数可能已经偏离全局模型很远,学术界把这个现象叫做client drift(客户端漂移)。FedAvg 把所有漂移后的模型做平均,得到的全局模型相当于“被各方扯向不同方向后的合力”,如果漂移太严重,这个合力就失去意义,模型甚至在某些类别上退化。
我实测过一个用 CIFAR-10 划分成严格 Non-IID(每个客户端只保留 2 类数据)的实验,FedAvg 在 100 轮时的测试精度比 IID 设置下低了将近 15 个点,而且训练过程非常震荡。这就是为什么后面需要 FedProx 和 SCAFFOLD 这些改进算法。
3. FedProx:给客户端加一道“近端约束”
3.1 近端项到底约束了什么
FedProx 是 MIT 的 Li 等人在 2020 年提出来的,论文名字叫《Federated Optimization in Heterogeneous Networks》。它的出发点很直接:既然客户端漂移是因为本地训练“太放飞自我”,那就在每个客户端的本地目标函数里加一个正则项,把它拉回全局模型附近。
每个客户端在本地求解的目标函数变成:
min_w F_k(w) + (μ/2) ||w - w_t||^2
其中 F_k(w) 是第 k 个客户端原本的本地损失函数,w_t 是第 t 轮的全局模型,μ 是一个非负超参数。第二项就是“近端项”(proximal term),它惩罚本地模型参数与全局模型参数的欧氏距离。当 μ=0 时,FedProx 直接退化为 FedAvg。
用生活类比就是:以前学生在各自教室里随便怎么学都行,现在学校规定,每次考试前必须先复习“全球统一课本”的前半部分,偏离太多了就要扣分。这个约束不是硬性禁止本地探索,而是让本地更新别跑太远,在“个性化”和“全局一致性”之间找平衡。
FedProx 还有一个很实际的设计:允许不同客户端使用不同的本地迭代轮数。真实场景下各设备算力差距很大,有些手机可能 5 秒就能完成 10 个 epoch,有些旧设备 10 秒只能跑 3 个 epoch。FedProx 允许这些“只完成一部分本地训练”的客户端也参与聚合,只要它们满足一定的近似度条件,这比 FedAvg 要求所有客户端都跑完固定轮次更贴近现实。
3.2 实施细节与超参选择
在代码层面,把 FedAvg 改成 FedProx 其实只差一行。以 PyTorch 为例,普通 FedAvg 的损失函数是 cross_entropy(output, target),FedProx 只需要在总损失后面额外加上近端项:
prox_term = (mu / 2) * sum((p - p_global).pow(2).sum() for p in model.parameters()) loss = criterion(output, target) + prox_term这里 p_global 是从服务器下载的全局模型参数,在本地优化过程中必须保持固定,可以用 detach() 或单独深拷贝一份。
关于 μ 的取值,我踩过不少坑。太小(比如 1e-5)接近 FedAvg,约束作用不明显;太大(比如 10)会让本地模型几乎复制全局模型,失去本地学习能力。我建议从 μ=0.01 开始,在验证集上观察收敛曲线,再调整数量级。在我做过的几个图像分类和推荐系统实验里,μ 在 0.01 到 1 之间通常能找到较好设置,而且数据异质性越强,可以尝试的 μ 就越大。
实施 FedProx 时还有一个细节容易被忽略:近端项里的 w_t 必须用“本轮下发时”的全局模型参数,而不是本地训练中途更新后的参数。如果你在动态图框架里不小心让全局模型参数被原地更新,近端项就失去了锚定作用,甚至会导致训练发散。
4. SCAFFOLD:用控制变量纠正更新方向
4.1 从“刹车”到“纠偏”
FedProx 的思路是给本地训练加约束,但它的上限受限于“约束强度”与“学习能力”之间的平衡。SCAFFOLD 换了一条路:不直接限制本地模型参数变化,而是通过估计并修正梯度方向,从根源上降低客户端漂移。
SCAFFOLD 全称是 Stochastic Controlled Averaging for Federated Learning,由 Karimireddy 等人在 2020 年提出。核心思想来自统计学里的“控制变量法”(control variates):既然本地梯度会偏离全局梯度,那我们就在本地更新中同时减去这种偏差,让每个客户端每一步都朝着更接近全局方向更新。
具体来说,服务器维护一个全局控制变量 c,表示全局梯度的估计方向;每个客户端维护一个本地控制变量 c_i,表示自己本地梯度的估计方向。在每个本地训练步中,客户端不再直接用本地梯度 g_i(x) 更新模型,而是用修正后的梯度:
g_i(x) - c_i + c
这个式子非常漂亮:减掉自己的偏差 c_i,加回全局方向 c,相当于把本地梯度“搬”到全局坐标系下再使用。如果某个客户端的数据分布特别偏,它的 c_i 会很大,正好抵消掉它本地梯度里“偏”的成分。相比 FedProx 的“软约束”,SCAFFOLD 是在每一步优化时做方向校正,所以理论上能更彻底地消除客户端漂移。
4.2 本地更新与聚合流程
SCAFFOLD 的一轮完整流程大概是:
- 服务器下发全局模型 x 和全局控制变量 c 给所有参与客户端。
- 每个客户端复制本地模型 x_i = x,同时保留自己的本地控制变量 c_i。
- 客户端在本地执行 S 步修正梯度下降:x_i ← x_i - η_l (g_i(x_i) - c_i + c)。
- 本地训练结束后,客户端计算模型参数增量 Δx_i = x_i - x,并更新本地控制变量,返回 Δx_i 和控制变量增量 Δc_i。
- 服务器用所有客户端的 Δx_i 更新全局模型 x,同时用所有客户端的 Δc_i 平均更新全局控制变量 c。
相比 FedAvg,SCAFFOLD 的聚合公式多了一个控制变量的更新环节。这个环节的意义是:每次训练后,全局控制变量会向“真实全局梯度方向”靠近一点,下一轮下发时对本地梯度的校正就更准确。理论分析也表明,在高度 Non-IID 环境下,SCAFFOLD 的收敛速度对数据异质性的敏感度远低于 FedAvg 和 FedProx。
顺带说一个容易被忽略的点:SCAFFOLD 虽然训练过程更稳,但代价是通信量大概翻倍——每个参与方除了要传模型参数增量,还要传一个控制变量增量。当模型本身很小(比如逻辑回归)时这个代价可以忽略,但如果是大模型,控制变量的通信成本就非常可观。
4.3 收敛快、通信翻倍的真实代价
那么,SCAFFOLD 什么时候值得用?
我个人的判断标准是:数据异质性严重(比如客户端之间标签分布几乎不重叠)、客户端参与率低(每轮只采样很小一部分),并且模型参数规模可控时,SCAFFOLD 的收益会非常明显。如果数据分布基本 IID,FedAvg 就已经很好,强制上 SCAFFOLD 只会增加工程复杂度和通信压力。
在一组实验里,我用 1000 个客户端、每轮采样 50 个、每个客户端只保留某一类样本的极端情景下,SCAFFOLD 达到目标精度的轮数大约是 FedAvg 的 1/3 左右。但换到 IID 数据后,差距迅速缩小到 10% 以内。这说明 SCAFFOLD 不是“全面升级版”,而是“特定条件下的性能增强版”。
5. 三种算法选型对比:别再只会 FedAvg 了
5.1 同一实验条件下的对比表
我提供一个自己常用的对比维度框架,你可以直接拿去评估项目选型:
| 维度 | FedAvg | FedProx | SCAFFOLD |
|---|---|---|---|
| 核心机制 | 参数加权平均 | 近端正则约束 | 控制变量纠偏 |
| 对 Non-IID 鲁棒性 | 低 | 中 | 高 |
| 计算复杂度 | 低 | 低(多一项正则) | 中(需维护控制变量) |
| 通信开销 | 1 倍 | 约 1 倍 | 约 2 倍 |
| 实现难度 | 低 | 低 | 中高 |
| 对异构算力的容忍度 | 低 | 高 | 中 |
| 超参数 | 学习率、本地 epoch 等 | 额外有 μ | 无额外关键超参 |
| 典型场景 | 数据分布均匀、设备能力接近 | 数据有一定偏移、设备算力差异大 | 数据高度 Non-IID、追求收敛速度 |
这个表里最容易被误解的是“通信开销”。FedProx 只改本地损失函数,上传下载的还是模型参数,所以通信量基本不变;SCAFFOLD 需要额外传控制变量,所以通信量翻倍。在做大规模系统设计时,这个差异可能比算法精度更重要。
5.2 根据项目情况选择聚合算法
如果你是项目决策者,我建议按下面这个逻辑来选:
- 先跑一版 FedAvg 作为 baseline,确认数据异质性程度。如果训练曲线平滑、精度达到预期,不建议折腾更多算法。
- 如果发现 FedAvg 在 Non-IID 下收敛慢甚至不收敛,先尝试 FedProx,改动最小、风险最低,通常能解决 60% 的问题。
- 如果 FedProx 仍然收敛太慢或者精度差强人意,同时模型参数体积可控、带宽预算充足,再考虑 SCAFFOLD。
- 如果项目里设备算力差异特别大,有些客户端只能做很少的本地迭代,FedProx 的灵活性优势会非常突出,但要注意 μ 和参与策略的配合。
另外别忽略工程层面的因素。FedAvg 和 FedProx 对任何常规深度学习框架都友好,SCAFFOLD 需要额外维护两组控制变量,在分布式调度、断点续跑、容灾恢复时都要额外处理。小团队从 FedProx 起步,往往比直接上 SCAFFOLD 更划算。
6. 实操踩坑记录与避坑清单
6.1 模拟 Non-IID 数据:Dirichlet 划分与参数选择
很多人在本地复现这三种算法时,第一个坑就是“数据划分方式不对”。如果你用随机划分(IID)跑出来的结果测 Non-IID 场景,当然看不出算法差异。
我常用的方式是 Dirichlet 分布划分。给定样本总数和客户端数量,先从 Dirichlet 分布中采样一个类别分布向量 p_k ~ Dirichlet(α),再把样本按这个分布分配到各客户端。α 越小,各客户端的类别分布越极端,大概取 α=0.1 时已经是非常强的 Non-IID 场景,而 α=1.0 时接近均匀分布。
还有一个细节:数据划分的随机种子必须固定。跨算法对比时,用的应该是“同一份划分结果”,而不是重新随机一遍,否则精度差异会混杂进划分随机性。我在代码里一般会把划分结果保存成文件,所有算法复用同一份。
6.2 训练过程中的常见问题速查
我在实际跑实验时遇到过不少典型问题,整理成速查表:
| 问题 | 可能原因 | 解决办法 |
|---|---|---|
| FedAvg 训练 Loss 震荡不收敛 | Non-IID 程度高、本地 epoch 过多、学习率偏大 | 换 FedProx/SCAFFOLD,降低本地 epoch 到 2~5,调小学习率 |
| FedProx 加了近端项后精度不升反降 | μ 设太大或太小 | 对数域搜索 μ,从 0.001 到 1,用验证集选最优 |
| 近端项没有生效 | 全局参数 w_t 被原地更新,锚点丢失 | 对全局参数 detach 或深拷贝,锁定初始 w_t |
| SCAFFOLD 通信量比预期高 | 控制变量和模型参数同时传输 | 压缩上传增量,或考虑混合精度量化 |
| 参与客户端过少导致波动大 | 每轮采样比例太低 | 提高参与比例,或增加评估轮次取滑动平均 |
| 本地控制变量初始化不合理 | 全部初始化为 0 但数据偏移大 | 先用 FedAvg 预热若干轮,再启动 SCAFFOLD |
6.3 我的一些实证经验
最后分享几点不一定写在论文里、但真实影响项目效果的经验。
第一,别迷信“更先进的算法”。我见过不少团队把 SCAFFOLD 当作默认聚合算法,结果因为通信开销太大、系统复杂度上升,整体表现反而不如简单调好学习率的 FedAvg。算法没有银弹,baseline 先跑好,改进才有意义。
第二,聚合算法的超参数要放在“联邦系统”整体里调。很多人只调本地学习率和 batch size,忽略了采样比例、参与客户端数量、本地 epoch 这些联邦特有的变量。我自己的经验是,先固定采样比例和本地 epoch,再调聚合算法特定参数,最后回头统一调学习率,这个顺序能省下很多时间。
第三,评估指标不要只用精度。在真实联邦场景里,通信轮数是比精度更昂贵的资源。我会在实验报告同时记录“达到目标精度所需轮数”和“累计通信量”两个指标,很多时候能帮你发现算法在成本上的本质区别。
第四,控制变量的初始化值得单独验证。SCAFFOLD 的收敛证明通常假设参与者足够多、控制变量初值接近真实梯度。如果项目冷启动时数据分布极端,先用 FedAvg 训练 10~20 轮收敛到一个可用模型,再切换到 SCAFFOLD,往往比从头就在 SCAFFOLD 上硬跑稳定得多。
最后,联邦聚合不是独立的算法模块,它和你用的加密协议、通信协议、客户端调度策略深度耦合。比如 secure aggregation 会引入额外噪声,这时 FedAvg 对噪声更敏感,而 SCAFFOLD 的方差校正机制反而能抵消一部分噪声影响。在做系统设计时,一定要把聚合算法放在完整的链路里评估。
这些经验都是我实际跑过、踩过之后总结出来的。你直接拿去用,至少能少走几周弯路。如果后面有机会,我再写一篇关于如何在生产系统里做聚合算法 A/B 测试的文章。