FedAvg论文学习(一)
前面跑 MNIST 的 MLP 和 CNN 时,我一直是在一台机器上训练:训练数据放在一起,模型从这些数据里算 loss、反向传播、更新参数。这个流程已经能看懂一点,但它有一个很容易被忽略的前提,就是数据可以集中起来。
读 FedAvg 这篇论文时,我才开始认真看另一个问题:如果数据本来就在很多不同设备上,模型还怎么一起训练?这篇文章不是 FedAvg 的复现实验记录。下面提到的实验数字都来自原论文,我目前还没有在本地把它复现出来。
论文是 McMahan 等人在 AISTATS 2017 发表的《Communication-Efficient Learning of Deep Networks from Decentralized Data》。Federated Averaging(联邦平均,FedAvg)是它提出并系统实验的一种联邦学习训练方法。
为什么现在读这篇论文
MLP 和 CNN 的训练让我先把“数据 -> 模型 -> loss -> 反向传播 -> 参数更新”这条线串起来了。但它们都属于单机训练:数据集中,模型也集中。
联邦学习,Federated Learning,FL,改变的是数据在哪里这个前提。原始训练数据可以留在客户端,服务器维护一个大家共同训练的全局模型。客户端收到模型后,在自己的数据上训练,再把模型更新传回服务器。
所以我想先弄明白的不是某一段 FedAvg 代码,而是几个更基础的问题:服务器和客户端之间到底来回传什么?为什么不每算一步就通信?多个客户端训练出的结果为什么不能随便做一个普通平均?
论文想解决什么问题
论文关注手机、平板等终端设备上的数据。用户输入的文字、照片、语音和使用记录可能留在设备上;这些数据既多,也不适合为了训练而全部集中上传。
联邦学习的基本想法是,原始数据尽量留在客户端,服务器只负责协调模型训练。它减少了集中收集原始数据的需要,但论文真正想优化的重点是:在数据分散、通信昂贵的情况下,怎样把模型训练出来,并尽量减少通信轮次。
我一开始容易把它理解成“数据不上传,所以已经安全了”。现在我觉得这个说法太快了。数据不集中确实减少了一类风险,但模型更新、通信链路和参与方本身仍然会带来新的问题。隐私边界放到后面再说,先把训练流程看清楚。
联邦优化和普通分布式训练有什么不同
论文把这里的优化问题称为 Federated Optimization(联邦优化)。它和数据中心里常见的分布式训练不太一样。
Non-IID(非独立同分布):不同客户端的数据分布可能不同。比如有的用户经常输入技术词汇,有的用户主要输入日常聊天;即使都是 MNIST,不同客户端看到的数字类别和比例也可能不一样。
Unbalanced(数据不均衡):不同客户端的数据量可能差很多。一个频繁使用应用的用户可能有大量样本,另一个用户可能只有少量样本。聚合时不能假装它们贡献的数据一样多。
Massively distributed(大规模分散):真实场景中的客户端可能非常多,而且并不会同时在线。每一轮都要求所有设备参加,基本不现实。
Limited communication(通信受限):终端设备会离线,网络速度和流量成本也有限。和数据中心相比,联邦学习里往往是通信更稀缺。
这四点让我理解了论文的主线:既然通信相对昂贵,就让客户端在两次通信之间多做一些本地训练,用本地计算换取更少的通信轮次。
一轮 FedAvg 到底发生了什么
论文用同步的通信轮次,communication round,组织训练。一轮 FedAvg 可以先按下面六步理解:
- 服务器保存当前的全局模型。
- 服务器随机选择一部分客户端参与这一轮。
- 服务器把同一份全局模型发给这些客户端。
- 每个客户端使用自己的本地数据训练若干轮。
- 客户端把本地训练后的模型参数或更新传回服务器。
- 服务器按各客户端的本地样本数加权平均,得到下一轮的全局模型。
这个“通信轮次”和单机训练里常说的 epoch 不是一回事。一次通信轮次包含了模型下发、本地训练、上传和服务器聚合;而 epoch 是客户端在自己本地数据上遍历训练集的次数。
FedSGD 和 FedAvg 的关系
论文不是凭空跳到 FedAvg,而是先从 Federated Stochastic Gradient Descent(联邦随机梯度下降,FedSGD)说起。整体训练目标可以写成:
其中,$K$ 是客户端总数,$n_k$ 是第 $k$ 个客户端拥有的样本数,$n$ 是全部样本数,$F_k(w)$ 是模型参数 $w$ 在第 $k$ 个客户端本地数据上的平均损失。这个式子说明,全局目标本身就应该考虑不同客户端的数据量。
FedSGD 可以理解为每个客户端根据当前全局模型计算一次本地更新,再让服务器聚合。FedAvg 则让客户端先在本地做多次小批量 SGD,再把训练后的模型交给服务器。对于本轮参与的客户端集合 $S_t$,加权聚合可以写成:
这里的 $w_{t+1}^{k}$ 是第 $k$ 个客户端本地训练后的模型参数。它不是简单的算术平均,而是让样本更多的客户端拥有更大的聚合权重。
论文把 $B=\infty$、$E=1$ 视作 FedSGD 的一个端点:整个本地数据作为一个 batch,只进行一次本地训练。减小 $B$ 或增加 $E$ 后,每次通信之间就会发生更多本地更新,这才是 FedAvg 最有代表性的地方。
三个关键参数:C、E、B
论文用 $C$、$E$、$B$ 控制每一轮有谁参与,以及每个客户端训练多少。
| 参数 | 论文中的含义 | 我现在的直观理解 | 需要注意的地方 |
|---|---|---|---|
C |
每轮选中的客户端比例 | 这一轮叫多少客户端参加 | 更大不一定线性减少通信轮数 |
E |
本地训练的 epoch 数 | 客户端收到模型后在本地练几遍 | 太大时可能出现平台期或发散 |
B |
本地 minibatch 大小 | 一次拿多少本地样本更新 | 越小通常意味着每轮有更多本地更新 |
对拥有 $n_k$ 个样本的客户端来说,本地更新次数大致和 $E n_k / B$ 有关。论文关心的正是:增加这些本地更新后,能不能显著减少必须和服务器通信的次数。
论文怎样构造 IID 和 Non-IID 数据
为了观察数据分布的影响,论文在 MNIST 上构造了两种客户端划分。
IID,Independent and Identically Distributed,独立同分布,设置里,作者先打乱 60,000 个训练样本,再平均分给 100 个客户端,每个客户端得到 600 个样本。这样每个客户端的数据分布都比较接近完整 MNIST。
pathological Non-IID,极端非独立同分布,设置则更刻意一些:先按数字标签排序,再切成 200 个 shard,每个 shard 有 300 个样本,最后给每个客户端随机分配 2 个 shard。这样大多数客户端只会看到大约两类数字。
这个例子让我把 Non-IID 看得更具体了。不过它更像一个压力测试,并不等于现实里每个用户都只会有两个类别的数据。真实差异还可能来自样本数量、类别比例、使用习惯和时间变化。
原论文的结果告诉了我什么
论文测试了 MNIST、CIFAR-10、Shakespeare 等任务。对我现在来说,最容易衔接的是 MNIST CNN。下面数据来自原论文第 6 页 Table 2,目标都是达到 99% 测试准确率:
| 任务与目标 | 方法与参数 | IID 通信轮次 | Non-IID 通信轮次 |
|---|---|---|---|
| MNIST CNN,99% accuracy | FedSGD:E=1, B=∞ |
626 | 483 |
| MNIST CNN,99% accuracy | FedAvg:E=20, B=10 |
18 | 173 |
在这组论文设置中,IID 条件下,通信轮次从 626 降到 18,大约少了 34.8 倍;极端 Non-IID 条件下,从 483 降到 173,大约少了 2.8 倍。
这组对比很直接地支持了论文的核心想法:客户端在本地多做一些训练,确实可能减少达到目标准确率所需的通信轮次。但 Non-IID 条件下的改善小得多,也说明客户端各自朝本地数据更合适的方向走时,服务器聚合会更难协调。
这里还要分清一件事:FedAvg 减少的是通信轮次,而它依靠的是更多客户端本地训练。通信变少,不代表总计算量、总耗时和设备能耗一定同时降低。
本地训练不是越多越好
看到 $E$ 变大能减少通信轮次后,很自然会问:那是不是每个客户端本地训练越久越好?论文给出的答案不是简单的“是”。
在 Shakespeare LSTM 实验中,本地 epoch 很大时,训练可能进入平台期,甚至发散;在论文的 MNIST CNN 设置里,较大的 $E$ 又没有表现出同样明显的退化。这说明 $E$ 的效果还和模型、学习率、数据分布以及训练阶段有关。
特别是在 Non-IID 数据上,客户端训练太久,可能越来越适应自己的局部数据,和其他客户端的更新方向差得更远。后来的研究会把这种现象叫作 client drift,客户端漂移;我现在先记住它的直觉,不把它当成已经完全掌握的概念。
数据不上传,不等于已经安全
这篇论文的出发点里确实有隐私考虑:原始数据不必为了训练长期集中到服务器。相比把所有数据都收集起来,这当然减少了直接暴露的范围。
但这不是“联邦学习天然安全”的证明。论文也提到,模型更新本身可能包含信息。服务器仍然会收到参数或梯度形式的更新,系统还可能面对恶意客户端、推断攻击、掉线和通信过程泄露等问题。
这和我之前整理的隐私保护概念能接上:FedAvg 先解决的是“分散的数据怎样共同训练”,差分隐私、SMPC、安全聚合和同态加密则从不同角度补充“训练过程中怎样少泄露信息”。它们有关联,但不能互相替代。
我现在理解了什么
读之前,我把 FedAvg 简化成“把多个客户端模型平均”。现在我觉得这句话不算错,但不够完整。
FedAvg 不是一种新的网络结构,模型仍然可以是 MLP、CNN 或 LSTM;损失函数、梯度下降和参数更新也仍然存在。它变化的是训练怎么组织:从许多客户端中抽样,让它们从同一个全局模型开始本地训练,再由服务器按样本数加权聚合。
我目前能用一句话概括它:FedAvg 让客户端在两次通信之间多做一些本地计算,用这些计算换取更少的通信轮次。论文展示了这种方法在特定 IID、Non-IID 和数据不均衡设置下的可行性,但不代表所有模型和真实设备环境都会得到一样的效果。
还没解决的问题和下一步
这次读论文让我先看懂了主流程,但我还没有自己实现服务器、模拟客户端、本地训练和加权聚合;也没有亲自观察 $C$、$E$、$B$ 以及 IID、Non-IID 划分会怎样影响曲线。
所以下一步会基于已经跑通的 MNIST MLP 和 CNN,做一个简化的 FedAvg 复现:把数据划成多个模拟客户端,再实现本地训练和服务器加权平均。到那时,才能把论文里的通信轮次、数据划分和实际训练曲线对起来。
论文信息
- 论文:Communication-Efficient Learning of Deep Networks from Decentralized Data
- 作者:H. Brendan McMahan、Eider Moore、Daniel Ramage、Seth Hampson、Blaise Agüera y Arcas
- 会议:AISTATS 2017
- arXiv:1602.05629v4
- 本文主要参考:联邦优化特点、FedAvg 算法、MNIST 数据划分、Table 2、本地 epoch 讨论和隐私方向





