pytorch_federated_learning
PyTorch 实现的 FedAvg/FedProx/SCAFFOLD/FedNova 联邦学习基线库,数据不出本地、模型参数协同训练
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
PyTorch 实现的 FedAvg/FedProx/SCAFFOLD/FedNova 联邦学习基线库,数据不出本地、模型参数协同训练
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象一个这样的场景: 你是一家医院的数据科学团队,手里积累了大量宝贵的医学影像数据,团队希望用这些数据训练一个更精准的 AI 诊断模型。但问题来了——这些数据涉及患者隐私,医院之间不能直接共享数据,否则会违反《个人信息保护法》和《数据安全法》。你只能眼睁睁看着隔壁医院的同类数据白白浪费。
这正是联邦学习(Federated Learning)要解决的问题。PyTorch Federated Learning 是 Rui Song 等研究者开源的联邦学习基线实现库,它让研究者和工程师无需从零搭建,就能快速跑通 FedAvg、FedNova、FedProx、SCAFFOLD 等主流联邦学习算法,在保护数据隐私的前提下实现跨节点协作训练。

图1:各联邦学习基线算法在 MNIST 数据集上的收敛曲线对比
传统机器学习要求把所有数据集中到一台服务器上训练——这在数据隐私日益敏感的今天变得越来越困难。联邦学习的核心思想是:让数据留在本地,只有模型参数在节点之间流转。 举个例子:假设有 100 家医院参与联邦学习,每家医院用自己的患者数据训练本地模型,然后只将模型梯度上传到中央服务器;服务器聚合所有节点的更新后,再将新的全局模型下发回各节点。如此循环迭代,最终所有医院都能受益于全局模型的提升,而原始数据从未离开本地。
这项技术最早由 Google 在 2016 年提出(McMahan 等,AISTATS 2017),并在 Gboard 键盘预测任务中落地。随着隐私法规日趋严格,联邦学习迅速扩展到金融(反欺诈模型协作)、医疗(跨医院影像诊断)、移动端(个性化推荐)等多个领域,成为隐私计算领域最受关注的技术方向之一。
本项目的代码组织非常清晰,采用了典型的客户端-服务器双层架构,核心模块如下:
| 模块 | 路径 | 功能 |
|---|---|---|
| 客户端基类 | fed_baselines/client_base.py | FedAvg 标准客户端 |
| 服务器基类 | fed_baselines/server_base.py | 中心化聚合服务器 |
| FedProx 客户端 | fed_baselines/client_fedprox.py | 加入近端项修正梯度 |
| SCAFFOLD 客户端 | fed_baselines/client_scaffold.py | 引入控制变量纠正偏差 |
| FedNova 客户端 | fed_baselines/client_fednova.py | 归一化梯度加权聚合 |
| 数据预处理 | preprocessing/baselines_dataloader.py | 自动下载 + Non-IID 划分 |
| 模型库 | utils/models.py | LeNet、AlexNet、ResNet、VGG 等 |
| 工具函数 | utils/fed_utils.py | 数据集参数映射、模型初始化 |
客户端逻辑(以 client_base.py 为例)包含以下步骤:
服务器逻辑(server_base.py)包含:
# 客户端本地训练核心(伪代码)
def local_update(self, model_state_dict):
self.model.load_state_dict(model_state_dict)
for epoch in range(self._epoch):
for batch in self.trainset:
loss = self.model(batch)
loss.backward()
self.optimizer.step()
return self.model.state_dict()
联邦学习算法(4 种):
支持的数据集(5 种): MNIST、Fashion-MNIST、EMNIST、SVHN、CIFAR-10、CIFAR-100,自动下载并按 Non-IID(标签分布非独立同分布)方式划分给各客户端——这更接近真实场景(各家医院数据分布差异大)。
本项目的定位是研究工具包,面向的是有 PyTorch 基础的算法研究者。它有明确的硬件要求:
torch.cuda.is_available() 是硬编码检查):每个节点需要在前向+反向传播中做大量矩阵运算,CPU 训练极慢使用方式非常直接——编辑 YAML 配置文件指定算法、数据集、模型和超参数,运行 python fl_main.py --config "./config/test_config.yaml" 即可。实验结果通过 postprocessing/eval_main.py 自动绘图展示收敛曲线。
值得指出的是,这个项目并非开箱即用的产品级系统,研究者在使用时需要注意:
gpu = 0,多卡环境需要手动修改联邦学习正在从学术研究走向工业落地。2020 年以来,蚂蚁集团"摩斯"平台、华为云隐私计算引擎、微众银行 FATE 框架相继推出,使联邦学习在金融风控和医疗 AI 领域逐步规模化。本项目作为学术基线库,代码简洁、逻辑清晰,非常适合作为理解联邦学习核心思想的入门材料——既不比 PyTorch 官方示例复杂太多,又能让你亲手体验 Non-IID 数据划分的挑战。
如果你想快速验证一个联邦学习想法,或在论文中对比基线算法性能,这个工具箱是值得一试的选择。