联邦聚合这个话题,在隐私计算和分布式深度学习圈子里已经不算新鲜了。但奇怪的是,真正敢说自己把FedAvg、FedProx、SCAFFOLD这三个算法用明白的人,我接触下来其实并不多。原因也很简单:大部分人看论文只看了公式,没把算法放进真实的数据分布、通信环境和客户端异构情况里去验证,一到项目落地就翻车。
我前后在几个联邦学习项目里把这三种聚合算法从零到一撸过完整流程,从最早的FedAvg打基线,到后来为了解决非独立同分布数据导致模型不收敛的问题改用FedProx和SCAFFOLD,中间踩过的坑确实不少。这篇文章不打算做教科书式的理论复述,我想用人话把三种算法的原理、公式、关键参数、适用场景和调试经验一次说清楚。如果你正在做联邦学习方案的选型、打算自己实现聚合算法,或者准备复现论文但被各种符号绕晕了,这篇内容应该能帮你少走不少弯路。
1. 联邦聚合问题,到底难在哪里
1.1 聚合在联邦学习流程里的“咽喉”位置
要理解为什么聚合算法这么重要,先得把联邦学习的完整闭环看清楚。整个训练流程大概是这样的:服务器先把当前全局模型参数下发到一批参与训练的客户端设备上;各客户端利用本地数据在本地模型基础上做几轮梯度下降;训练完成后,客户端把模型更新或模型参数上传回服务器;服务器对收集到的更新做聚合,生成新的全局模型;然后进入下一轮循环。
这个过程里,聚合是唯一一个能把所有客户端“意见”汇总起来的环节。本地训练只负责让每个客户端在本地数据上变聪明,但全局模型能不能兼顾所有客户端的数据分布,完全取决于聚合这一步怎么设计。如果把全局模型比作一个班级的最终成绩单,本地训练就是每个学生在自己家里刷题,聚合则是老师把大家的错题本汇总成一套新卷子的过程。错题本怎么合并、每本错题集权重多大、要不要剔除明显离谱的答案,这些就是聚合算法要解决的问题。
1.2 非独立同分布数据是怎样“搅局”的
理想情况下,如果所有客户端的数据分布都独立同分布,那联邦学习的问题会简单很多。但现实里,各客户端的数据几乎一定是非独立同分布的——不同用户的照片风格不同,不同医院的病例分布不同,不同城市的消费行为差异巨大。这种数据异构性会导致一个非常棘手的现象:各客户端在本地数据上训练出来的模型,很可能朝完全不同的方向漂移。
我打一个不太严谨但很直观的比方。三个老师分别辅导三个不同水平的学生,老师A面对的全是基础薄弱的学生,所以他总结出“先补基础”的教法;老师B面对的全是拔尖学生,他主张“直接上难题”;老师C面对的是中等生,他觉得“多刷中等题”最有效。这三个老师各自把自己的教学心得写成一本书,然后把三本书的内容逐字平均成一本新教材。这本新教材大概率既不适应基础差的学生,也不适应拔尖的学生,甚至连中等生都觉得别扭。
模型参数平均就是这个道理。当各客户端的本地模型差异过大时,简单平均会得到一个被“摊平”的模型,这个模型在任何单一客户端的数据上表现都不好,甚至在全局数据上也表现平庸。
1.3 三大算法其实是在解决不同层面的痛点
FedAvg、FedProx、SCAFFOLD三者经常被放在一起对比,但很多人误以为它们是“进阶版”关系,实际并非如此。它们各自切入点不同,解决的问题层面也不同。
| 算法 | 核心创新点 | 针对的主要痛点 |
|---|---|---|
| FedAvg | 按样本量加权的参数平均 | 建立联邦学习的基本聚合范式 |
| FedProx | 在本地目标中加入近端正则项 | 缓解数据异构导致的本地模型漂移 |
| SCAFFOLD | 引入控制变量(控制变分)修正本地梯度方向 | 从梯度估计偏差层面消除客户端漂移 |
FedAvg是地基,FedProx和SCAFFOLD是在地基建起来的两栋方向不同的房子。理解这个区别,等真正选型的时候才不会只看名字就下结论。
2. FedAvg——简单但绝不简陋的基线
2.1 FedAvg的核心流程和聚合公式
FedAvg最早由McMahan等人提出,核心思路一句话就能说完:各客户端用本地数据训练若干个epoch后,把更新后的模型参数(或参数增量)上传到服务器,服务器按每个客户端拥有的样本量在总样本量中的占比做加权平均,得到新的全局模型。
如果用公式表示,假设有K个客户端参与本轮聚合,第k个客户端拥有 n_k 个样本,总样本数 n = Σ n_k,本地训练后的模型参数为 w_k,那么全局模型更新为:
w_new = Σ_{k=1}^{K} (n_k / n) · w_k这个公式简单到看起来不像一个能发顶会的方法,但它在很多场景下确实有效,原因在于它抓住了联邦学习的本质:全局模型应该是所有客户端本地经验的加权平均,而样本量是衡量一个客户端数据“代表性”最直接、最公平的指标。
2.2 收敛理论给我们的启示:IID假设是关键前提
FedAvg的理论分析通常是在凸问题、所有客户端全量参与、本地更新步数一致的假设下证明收敛的。但实际应用中,这几个假设几乎不可能完全满足。尤其是IID假设,一旦客户端数据分布差异明显,FedAvg的收敛速度会急剧下降,甚至出现来回震荡、始终无法收敛的情况。
我自己测试过一组实验:用同一个模型结构、同样的超参数,在IID数据切分下FedAvg收敛得很漂亮,但改成按用户维度切分数据(模拟真实非IID场景)后,训练loss曲线开始出现明显波动,最终模型的测试准确率下降了将近8个百分点。这个数字不一定适用于所有任务,但趋势非常一致——数据异构越严重,FedAvg的表现下降越明显。
2.3 FedAvg会崩的三种典型场景
我把实战中遇到过的FedAvg翻车场景归纳成三类,你可以对照自己的项目排查。
第一种是客户端数据分布极端不均衡时。有的客户端可能有几千条样本,有的只有几十条,且两者的类别分布完全不同。这种极端情况下,按样本量加权平均后,小样本客户端的更新几乎被淹没,大样本客户端又在凭借本地偏差把全局模型往自己的方向带,结果就是全局模型越训越偏。
第二种是本地epoch设置过大时。当参与方本地训练轮数过多,模型在本地数据上过拟合严重,本地模型和初始全局模型的差距会非常大。这个时候做参数平均,等于把几个“偏见极深”的模型硬生生揉在一起,效果自然好不了。
第三种是客户端参与率低且参与客户端数据分布差异巨大时。每轮只有少量客户端参与,如果这些客户端本身数据分布就高度不一致,聚合出来的全局模型会随参与客户端的组合变化而剧烈波动,训练曲线像心电图一样上下乱跳。
2.4 什么时候仍然值得选FedAvg
别因为上面这些问题就一棍子打死FedAvg。实际项目中,它仍然是首选的基线算法和很多业务场景的正解。
如果你的数据分布经过预处理后已经比较接近IID,或者各客户端的数据分布大致均匀,FedAvg完全够用,而且它通信量最小、实现最简单、超参数最少,几乎不会因为聚合部分出幺蛾子。另外,如果你需要一个“对照组”来判断其他算法到底有没有效果,用FedAvg做基线是唯一合理的做法——不看基线就评估新算法,结论是站不住的。
3. FedProx——给本地更新套上“缰绳”
3.1 从FedAvg到FedProx:问题到底出在哪
FedProx的提出者意识到,FedAvg在非IID数据下表现不佳,一个很重要的原因是:本地训练过程没有“全局视野”,每个客户端只管在本地数据上努力优化本地目标,完全不在乎本地模型和全局模型之间已经拉开了多大距离。这种“各自为政”的训练方式,累积出来的偏差会在聚合时集中爆发。
FedProx的核心改动非常巧妙——它没有去改聚合方式,而是在每个客户端的本地优化目标里加了一个约束项,限制本地模型不要偏离初始全局模型太远。
3.2 近端正则项是怎么起作用的
FedProx的本地优化目标可以写成:
min_w [ F_k(w) + (μ / 2) · || w - w_t ||² ]其中 F_k(w) 是第k个客户端在本地数据上的原始损失函数,w_t 是服务器下发的当前全局模型参数,μ 是近端项的惩罚系数。这个二次正则项的直观含义是:你可以在本地数据上优化,但不要让参数跑离全局模型太远,否则会受到“惩罚”。
这个思路其实和机器学习里常见的正则化一脉相承。就像让一个员工在分店独立工作,但总部要求他定期汇报,并且每周的任务方向不允许偏离总公司战略太远。既给了本地灵活性,又控制了全局方向的一致性。
在实现上,FedProx只是改变了客户端本地训练时的损失函数,聚合部分仍然复用FedAvg的加权平均方式,因此工程改造量非常小。这也是它在工业界接受度高的原因之一。
3.3 μ到底怎么调:从0开始逐步试探
μ的选择是FedProx落地时最关键的调参环节,没有之一。μ=0时,FedProx就退化成FedAvg,所以调参的正确姿势是从0附近往大了试探。
我个人的经验是,先设μ=0.01跑一个短周期,观察训练loss和验证精度。如果loss曲线仍然剧烈震荡,说明约束力不够,把μ提高到0.1;如果曲线很平稳但收敛速度明显变慢,说明约束过强,回退到0.03附近再试。
| μ取值 | 典型现象 | 建议 |
|---|---|---|
| 0 | 等价于FedAvg,非IID下可能震荡 | 作为基线对比 |
| 0.01~0.03 | 约束温和,适合轻度数据异构 | 推荐的起始区间 |
| 0.1~0.3 | 约束较强,适合严重数据异构或客户端本地训练轮次较多 | 异构严重时使用 |
| >1 | 收敛缓慢,本地模型几乎没机会学习 | 一般不建议 |
需要注意的是,μ的理想取值和数据异构程度、本地epoch数量、模型结构都有关系。我自己在图像分类任务上试过,当本地epoch从5增加到20时,最优μ大约从0.01上升到0.1,趋势很明显——本地训练越“贪”,越需要更强的缰绳。
3.4 FedProx的“副产物”优势:容忍异构计算
FedProx论文里还有一个很实际的设计:不同客户端每轮本地训练的epoch数可以不相同,不再要求所有客户端做同样多的工作。这在真实场景中太重要了。
我在实际项目里遇到过大量异构设备,手机型号不同、算力不同、当前电量不同,如果强制所有客户端都做5个epoch,部分设备可能要跑很久,导致整轮训练被拖慢。FedProx允许这些设备只做2~3个epoch也能参与聚合,因为它通过近端项限制了它们对全局模型的影响,不会因为“做得少”而带来过大的偏差。
4. SCAFFOLD——用控制变量纠正客户端漂移
4.1 为什么FedProx还不够
FedProx的思路本质上是“被动防守”——它假设本地模型会漂移,所以加一个约束把它往回拽。但约束再强也只是在限制漂移的幅度,并没有真正解决漂移的方向问题。如果某个客户端的数据分布和全局分布差异很大,它的本地梯度方向本身就是偏的,正则项只能让它别跑太远,却改变不了它跑偏的方向。
SCAFFOLD走的是另一条路:它认为本地模型漂移的根源,是各客户端在本地数据上估计的梯度方向与全局真实梯度方向之间存在系统性偏差。既然知道偏差存在,那就干脆显式地估计这个偏差,并在更新时把它减掉。
4.2 控制变量的核心思想在联邦学习里的落地
SCAFFOLD引入了两个控制变量:服务器维护一个全局控制变量c,代表全局梯度的“基准方向”;每个客户端维护一个本地控制变量 c_k,代表该客户端本地梯度的“基准方向”。本地更新时,不再直接用本地梯度去更新模型,而是用修正后的梯度:
x_new = x - η · ( g_k(x) - c_k + c )这里 g_k(x) 是第k个客户端在本地数据上计算的梯度。当本地梯度方向和全局基准方向一致时,修正项抵消为零,客户端正常更新;当本地方向和全局方向偏离明显时,修正项会把更新拉回更接近全局方向的位置。
这个机制的直观理解,可以类比成一队人夜跑时各自拿着手电筒。每个人以为自己是在沿着正确的方向跑,但有人偏向左、有人偏向右。SCAFFOLD做的事是给每个人配一个GPS基准方向,告诉他们“你感觉的方向和全局方向的偏差是多少,请自动修正”。
值得注意的是,这种修正机制对客户端漂移的纠正是“主动”的——它改变的不仅是更新幅度,还有更新方向,这正是它和FedProx最本质的区别。
4.3 两个控制变量是怎么更新和维护的
SCAFFOLD的实际流程比FedProx稍微复杂一点,但理解之后实现起来并不困难。
服务器端维护全局参数x和全局控制变量c。每个客户端在本地维护c_k。每轮训练开始前,服务器把x和c同时下发给本轮选中的客户端。
客户端收到后,在本地执行多步更新,每一步都用修正后的梯度 x = x - η·(g_k(x) - c_k + c) 来迭代。本地迭代结束后,再根据本地模型的变化量、本地迭代步数、学习率等信息,计算新的本地控制变量c_k并上传。服务器聚合所有被选中客户端的更新量和控制变量变化量,更新全局参数x和全局控制变量c。
在论文的原始设定里,服务器更新全局控制变量c时,会往全局c里加一个平均的Δc_k。这个机制的核心作用,是用历史梯度信息给当前估计“去偏”,从而加快收敛速度。SCAFFOLD论文的一个重要结论是:在数据异构严重且所有客户端全量参与的情况下,SCAFFOLD的收敛速度不会随着数据异构程度的增加而恶化,而FedAvg会。
4.4 SCAFFOLD的代价:通信量翻倍和额外存储
SCAFFOLD这么好用,但并没有取代FedProx成为工业界主流,原因在于它的代价不小。最主要的问题在于通信开销翻倍——服务器不仅需要下发模型参数,还需要下发全局控制变量c,控制变量的维度和模型参数量完全一致,所以每轮通信的传输量大约是FedAvg的两倍。
对于一个上亿参数的大模型,这个额外开销是相当可观的。如果通信带宽本就不宽裕,增加的传输时间很可能抵消掉SCAFFOLD带来的收敛加速收益。此外,每个客户端还要额外维护c_k,多了一份和模型同维度的存储开销,对内存受限的边缘设备也不友好。
所以SCAFFOLD的适用场景很明确:数据异构程度确实很严重、通信不是最紧的瓶颈、客户端存储空间足够。如果你的项目符合这些条件,SCAFFOLD可能是比FedProx更值得尝试的选择。
5. 选型与实战:三种算法在项目里到底怎么落地
5.1 我的选型判断流程:三步走
每次接到联邦学习相关的需求,我习惯先按一个固定流程走,能省掉大量无意义的实验时间。
第一步,定量评估数据异构程度。可以先跑一个FedAvg小规模实验,观察训练loss的波动幅度;如果震荡明显,再用数据分布分析工具看各客户端类别分布差异。一个很实用的经验:客户端之间类别分布KL散度平均值超过一定阈值,基本可以判断为严重异构。
第二步,分析通信和计算约束。每轮通信的带宽上限、客户端设备是否异构、存储是否紧张,这些直接决定了能否考虑SCAFFOLD类似的“多传一份数据”的方案。带宽不足优先考虑FedProx,传一个标量超参数的成本几乎可以忽略。
第三步,定基线、定对比。不管最终用哪个算法,我都会先跑一版FedAvg作为基准,然后根据痛点切换到对应方案。这样做的好处是,你能清楚地看到新算法到底在哪个指标上带来了增益,判断这个增益是否值得工程复杂度的增加。
5.2 一组可落地的参数配置参考
这里给出一组我在图像分类任务上实测过的初始参数,数据集为CIFAR-10风格的彩色图像分类,客户端数量100个,每轮采样10个客户端。这组参数不一定适合你的任务,但作为起点很合理。
| 参数项 | FedAvg | FedProx | SCAFFOLD |
|---|---|---|---|
| 本地epoch | 5 | 5~10(可异构) | 5 |
| 本地batch size | 32 | 32 | 32 |
| 学习率 | 0.01 | 0.01 | 0.01 |
| 优化器 | SGD | SGD | SGD |
| 近端系数μ | 不适用 | 0.01起步,最大不超过0.3 | 不适用 |
| 客户端采样率 | 0.1 | 0.1 | 0.1(全量参与效果更好) |
| 通信轮次 | 100~200 | 100~200 | 50~150 |
SCAFFOLD在参与率低时效果会打折,因为它依赖控制变量的准确性,参与客户端越少,全局控制变量的估计噪声越大。如果实际场景没法做到较高的参与率,SCAFFOLD的优势会被明显削弱。
5.3 一个最简单的PyTorch聚合代码骨架
三种聚合算法的实现在聚合这一步区别很小,我经常把它们写在一个模块里方便切换。
import torch def fed_avg(global_model, client_models, client_sizes): total_size = sum(client_sizes) global_dict = global_model.state_dict() for key in global_dict: global_dict[key] = sum( client_models[i].state_dict()[key] * (client_sizes[i] / total_size) for i in range(len(client_models)) ) global_model.load_state_dict(global_dict) return global_model def fed_prox_aggregate(global_model, client_models, client_sizes): # FedProx的聚合部分和FedAvg一致,差别在客户端本地损失函数 return fed_avg(global_model, client_models, client_sizes) def scaffold_aggregate(global_x, global_c, client_delta_x, client_delta_c, client_sizes): total_size = sum(client_sizes) for i in range(len(client_delta_x)): weight = client_sizes[i] / total_size for key in global_x: global_x[key] += weight * client_delta_x[i][key] global_c[key] += weight * client_delta_c[i][key] return global_x, global_c这段代码把SCAFFOLD的聚合精简到了核心逻辑:全局模型和全局控制变量都按样本量加权更新。实际工程里还需要处理设备分发、容错、模型加密等一堆琐事,但聚合逻辑本身,就这么简单。
5.4 常见问题排查速查表
我在调试这三个算法的过程中,整理了一份高频问题速查表,每次遇到类似问题直接对症下药。
| 症状 | 可能原因 | 排查与解决方案 |
|---|---|---|
| 训练loss输出NaN | 学习率过高、特征数值异常、客户端上传了异常更新 | 降低学习率,增加梯度裁剪,在聚合前检查参数范数并剔除异常客户端 |
| 训练loss不降或震荡 | 数据异构严重,FedAvg失效 | 换成FedProx并调大μ,或改用SCAFFOLD |
| 验证精度先升后崩 | 本地epoch过大导致过拟合和漂移 | 减小本地epoch,FedProx提高μ,SCAFFOLD检查控制变量更新是否正常 |
| 聚合速度慢,训练迟迟不收敛 | 学习率过低或客户端参与率过高导致等待时间长 | 提高学习率,减少每轮参与客户端数,考虑异步聚合方案 |
| SCAFFOLD表现反而不如FedAvg | 参与率过低,控制变量估计噪声太大 | 提高参与率,或改用FedProx |
| 通信耗时成为瓶颈 | 模型过大,SCAFFOLD的额外控制变量通信加剧问题 | 考虑模型压缩、梯度量化,或退回FedProx |
最后分享点个人实际体会
这三个算法我反复用下来,最大的感受是:没有哪个算法是银弹,选型永远是在数据异构程度、通信成本、工程复杂度三者之间找平衡点。如果只给一条建议,我会说——先老老实实跑通FedAvg拿到基线,再根据loss曲线的具体形态决定下一步走向。震荡严重就试FedProx,收敛慢且通信允许就试SCAFFOLD,绝大多数项目走到这一步就已经能解决问题了。
另外在实现层面,不管用哪个聚合算法,都建议在聚合前后把全局模型的参数范数、梯度范数打点记录下来。很多联邦学习的问题不一定是聚合算法本身的锅,而是数据分布、超参数、设备故障共同作用的结果,没有这些日志,排查起来就像在黑暗中找钥匙。做好观测,再谈优化,这是我在这个领域踩过最多坑之后总结出的最重要一条经验。