Federated-Learning-PyTorch
PyTorch实现的经典联邦学习框架,完美复现FedAvg论文,支持IID/Non-IID数据分布实验
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
PyTorch实现的经典联邦学习框架,完美复现FedAvg论文,支持IID/Non-IID数据分布实验
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象一下这样一幅画面:三家三甲医院的影像科医生们,每天都在各自的医院里拍片、读片、积累经验。然而,当他们想合作训练一个更聪明的 AI 辅助诊断模型时,一个无法绕开的墙横亘在面前——患者隐私。HIPAA 法规、《个人信息保护法》明确规定,患者病历数据不能离开医院服务器,更别说上传到某个云端进行集中训练。
怎么办?2016 年,Google 的 H. Brendan McMahan 等人在论文《Communication-Efficient Learning of Deep Networks from Decentralized Data》中提出了一个优雅的解法——联邦学习(Federated Learning)。它的核心思想是:不要把数据搬到模型身边,而是把模型送到数据身边。
今天我们要分析的这个开源项目 Federated-Learning-PyTorch,正是对这篇经典论文的 PyTorch 实现,被全球研究者和学生广泛引用(GitHub 1438 ★,Fork 461 次)。
想象你有一群小学生(称为"客户端"),每人手里都有一叠写着不同动物图片和名字的卡片(各自的私有数据)。传统机器学习是让老师把所有卡片收上来,集中教一个 AI 认识所有动物。但联邦学习不是这样运作的——
老师先给每个小学生发一份空白答卷(初始模型),让他们各自在本地看自己手里的卡片学习几天。然后每个小学生把自己学到的知识精华(梯度/模型权重)交给老师,而不是把原始卡片交出去。老师把所有学生的精华汇总平均,形成新的、更聪明的答卷,再发给大家继续学习——如此循环多轮。
这样一来,学生的隐私卡片始终在自己手里,老师(中央服务器)永远不知道哪个学生看的是哪张卡片,但最终每个学生手里的 AI 都变得非常聪明。
项目代码结构清晰,划分为 6 个核心模块:
| 模块 | 文件 | 职责 |
|---|---|---|
| 配置层 | options.py | 统一管理所有超参数(联邦轮次、客户端数、本地学习率等) |
| 数据层 | sampling.py | 实现 IID 和 non-IID 两种数据分配策略 |
| 模型层 | models.py | 定义 MLP、CNNMnist、CNNFashion_Mnist、CNNCifar 四种神经网络 |
| 训练层 | update.py | 本地客户端训练逻辑、验证集切分、梯度更新 |
| 聚合层 | federated_main.py | 联邦聚合主循环:采样客户端 → 本地训练 → 权重平均 → 全局更新 |
| 工具层 | utils.py | 数据加载、FedAvg 聚合算法、日志输出 |
核心聚合算法 FedAvg(Federated Averaging) 在 utils.py 中实现:
def average_weights(w):
w_avg = copy.deepcopy(w[0])
for key in w_avg.keys():
for i in range(1, len(w)):
w_avg[key] += w[i][key]
w_avg[key] = torch.div(w_avg[key], len(w))
return w_avg
这段代码将各客户端返回的模型参数字典按元素逐个加权平均,生成新的全局模型参数。逻辑简洁,是整个联邦学习的精髓所在。
sampling.py 中实现了两种关键的数据划分策略,这是理解联邦学习性能差异的关键:
IID(独立同分布):将 MNIST 60000 张训练图均匀随机分配给 100 个客户端,每人 600 张。模拟理想情况,数据分布一致。
Non-IID(非独立同分布):将数据按标签排序后切分 shards,每个客户端只拿到 2 个 shards。例如张三只拿到了数字 0-1 的图片,李四只拿到了 2-3 的图片——这更贴近真实世界的场景:不同医院、不同科室的数据分布天然不同。
项目还实现了 non-IID unequal 模式,即各客户端数据量不均衡,模拟真实企业环境中的数据分布差异。
README 中提供了在 MNIST 数据集上的对照实验结果(10 轮,100 客户端,每轮采样 10%):
| 模型 | 集中训练准确率 | 联邦(IID) | 联邦(Non-IID) |
|---|---|---|---|
| MLP | 92.71% | 88.38% | 73.49% |
| CNN | 98.42% | 97.28% | 75.94% |
解读:
项目依赖简洁,对应 requirments.txt(注意文件名拼写):
pytorch=1.2.0 — 核心深度学习框架torchvision=0.4.0 — 数据集加载numpy=1.15.4 — 数值计算tensorboardx=1.4 — 训练可视化(TensorBoard 日志)matplotlib=3.0.1 — 训练曲线绘制tqdm — 进度条支持的模型和数据集组合:
| 数据集 | MLP | CNN |
|---|---|---|
| MNIST | ✅ | ✅(CNNMnist) |
| Fashion-MNIST | ✅ | ✅(CNNFashion_Mnist) |
| CIFAR-10 | ✅ | ✅(CNNCifar) |
尽管项目影响力巨大,作为学术研究基准实现,它也存在明显局限:
1. 通信效率被论文"超越"了:论文标题强调"Communication-Efficient",但本实现每轮传输完整模型参数(几十 MB 级别)。近年来 Google 提出的 Gradient Quantization、Sketching 等压缩技术,以及 FedProx、SCAFFOLD 等收敛优化方法均未包含。
2. 没有安全聚合(Secure Aggregation):参数以明文传输和聚合,在恶意服务器场景下存在信息泄露风险。
3. 依赖 PyTorch 1.2.0:该版本发布于 2019 年,存在大量已知安全漏洞,且不兼容新版 torchdistx 等优化库。
4. 仅支持横向联邦:不支持纵向联邦(特征维度不同)和联邦迁移学习等高级场景。
5. 缺乏差分隐私:没有内置 DP-SGD 或 PATE 等隐私保护机制。
该项目作为联邦学习领域的"Hello World",在学术和工业界都有深远影响。它催生了大量后续改进工作:
从 Google 在 Gboard 中首次将联邦学习落地,到苹果 CoreML 在 iOS 15 中支持 on-device 联邦学习,再到如今大模型时代的 FedML、FLUTE,联邦学习已经从学术概念演化为隐私计算的重要支柱。
# 安装依赖(注意是 requirments.txt 而非 requirements.txt)
pip install -r requirments.txt
# 基线对比实验(集中训练)
python src/baseline_main.py --model=mlp --dataset=mnist --epochs=10
# 联邦学习实验(IID)
python src/federated_main.py --model=cnn --dataset=mnist --iid=1 --epochs=10 --gpu=0
# 联邦学习实验(Non-IID)
python src/federated_main.py --model=cnn --dataset=cifar --iid=0 --epochs=10 --gpu=0
运行结果会自动保存到 save/ 目录,训练日志输出到 TensorBoard。