federated-learning
PyTorch 实现的经典联邦学习 FedAvg 算法,复现 Google 论文,支持 MNIST/
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
PyTorch 实现的经典联邦学习 FedAvg 算法,复现 Google 论文,支持 MNIST/
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
2016年,Google在一篇里程碑论文《Communication-Efficient Learning of Deep Networks from Decentralized Data》中首次提出**联邦学习(Federated Learning)**概念。这一范式诞生的背景极具现实意义:移动互联网时代,用户数据分散在数以亿计的终端设备上——手机输入法记住的词汇、智能手表记录的健康数据、银行的交易记录——这些数据蕴含巨大价值,却因隐私法规和数据安全顾虑无法集中到数据中心进行传统机器学习。
联邦学习的核心洞察:与其把数据汇聚到模型,不如让模型流动到数据所在之处。多个参与方(称为"客户端")在本地利用各自数据训练模型,仅将模型参数更新(而非原始数据)上传至中央服务器,由服务器聚合各方参数生成全局模型,再分发回各客户端迭代优化。这种"数据不动模型动"的机制,从根本上解决了数据隐私与协作学习之间的矛盾。
本项目 shaoxiongji/federated-learning 是这篇经典论文的 PyTorch 实现,由开发者 shaoxiongji 于2019年开源,目前累计获得 1516 Stars,是 GitHub 上最受欢迎的联邦学习入门级复现之一。项目被标记为 deep-learning、federated-learning、pytorch 三个主题,是 federated-learning 领域具有代表性的学习参考项目。
项目完整实现了 FedAvg(Federated Averaging) 算法,这是联邦学习领域最基础的聚合策略。其工作流程如下:
第1步:初始化全局模型 中央服务器初始化一个全局神经网络模型(如 CNN 或 MLP),并将初始参数下发给所有参与的客户端设备。
第2步:本地训练 每个客户端使用本地私有数据集对全局模型进行若干轮本地训练(Local Epochs)。以 MNIST 手写数字识别为例,假设有100个客户端,每个客户端获得约600条MNIST数据,各自独立运行 SGD 优化器更新本地模型权重。
第3步:参数上传
各客户端将本地训练后的模型参数(而非原始数据)上传到中央服务器。项目核心代码 models/Fed.py 中的 FedAvg 聚合逻辑极为简洁——对所有客户端的参数张量按索引逐层求均值:
def FedAvg(w):
w_avg = copy.deepcopy(w[0])
for k in w_avg.keys():
for i in range(1, len(w)):
w_avg[k] += w[i][k]
w_avg[k] = torch.div(w_avg[k], len(w))
return w_avg
第4步:全局聚合与分发 服务器将聚合后的全局参数回传给所有客户端,客户端用新参数替换本地模型,开始下一轮迭代。这一过程循环往复,直至全局模型收敛。
联邦学习区别于传统分布式学习的一个关键挑战是数据异质性(Data Heterogeneity)——各客户端的数据分布可能差异巨大(设备类型、用户习惯、地理位置等)。项目在 utils/sampling.py 中实现了两种数据划分策略:
IID(独立同分布):将数据集均匀随机分配给各客户端,每个客户端获得相似分布的数据子集。这相当于传统分布式学习的场景。
Non-IID(非独立同分布):这是更贴近现实的设置。项目采用了"标签分片"策略——先将 MNIST 按数字标签(0-9)排序,再依次分配给不同客户端,使每个客户端只持有部分类别的数据。例如,客户端A可能只见过数字0、1,而客户端B只见过数字8、9。这种极端的 Non-IID 设置会显著降低 FedAvg 的收敛速度,逼迫研究者寻找更鲁棒的聚合算法。
项目代码量约600行,结构清晰,适合作为联邦学习入门教材:
| 文件 | 功能 | 核心内容 |
|---|---|---|
main_fed.py | 联邦学习主流程 | 客户端采样、循环训练、FedAvg 聚合、精度测试 |
main_nn.py | 普通集中训练对照 | 基准训练流程,用于对比联邦与集中的性能差异 |
models/Nets.py | 神经网络定义 | MLP、CMNIST、CNNCifar 三种模型架构 |
models/Update.py | 本地更新逻辑 | LocalUpdate 类实现客户端本地训练循环 |
models/Fed.py | FedAvg 聚合 | 核心参数聚合算法 |
models/test.py | 测试评估 | 全局模型在测试集上的准确率计算 |
utils/sampling.py | 数据划分 | IID / Non-IID 数据集分配策略 |
utils/options.py | 命令行参数 | 超参数解析(epochs、客户端数、学习率等) |
关键超参数可通过命令行灵活配置:
python main_fed.py --epochs 10 --num_users 100 --frac 0.1 --local_ep 5 --dataset mnist --model cnn --iid
其中 --frac 0.1 表示每轮只随机选取10%的客户端参与训练,--local_ep 5 表示每轮本地训练5个 epoch。
GPU支持:代码中包含 CUDA 支持逻辑,当检测到可用 GPU 时自动将模型和数据迁移至 GPU 加速。但依赖 torch==0.4.1(发布于2018年的古老版本),该版本对应的 CUDA 版本较老,现代 GPU 可能存在兼容性问题。
计算量估算:以 MNIST 数据集、100个客户端、CNN模型为例,单轮联邦训练的计算量约为单个客户端本地训练100倍的全局开销——因为每个客户端都要跑完整的本地 epoch,再由服务器协调多轮通信。对于 CPU-only 环境,完整实验可能需要数小时。
资源消耗:RAM 需求约4GB(加载 MNIST/CIFAR10 数据集),磁盘约2GB(存放数据集缓存和训练曲线图)。
本项目的定位是教学级复现,而非生产级框架,存在以下局限:
不支持真实场景的隐私保护:模型参数在传输过程中以明文形式上传,存在梯度泄露(Gradient Leakage)风险。真实联邦学习系统需要差分隐私(Differential Privacy)或安全多方计算(SMPC)保护。
串行模拟,非真实分布式:代码在一个进程中顺序模拟多个客户端,并未真正连接多台物理设备。对于想要体验真实端-云联邦通信的研究者,需要 Flower、PySyft 等框架。
古老 PyTorch 版本:requirements.txt 指定 torch==0.4.1,与现代 PyTorch 2.x 存在 API 差异,直接 pip install 可能在最新 Python 环境(如 Python 3.12+)下安装失败。
CIFAR10 Non-IID 未实现:代码注释明确写明"only consider IID setting in CIFAR10",Non-IID 仅在 MNIST 上可用。
本项目虽然简洁,却是理解联邦学习的重要起点。2016年至今,联邦学习已从学术概念发展为 Google、Apple(输入法预测)、微众银行(风控模型)等企业的核心技术。项目作者在 README 中附上了 Zenodo DOI 链接,表明代码可被学术论文引用。
此后,联邦学习领域涌现了大量进阶框架:Google 的 TensorFlow Federated (TFF)、PySyft、Flower(支持任意 ML 框架的通用联邦框架)等。本项目可视为理解这些进阶工具的理论基础——只有亲手跑通 FedAvg 的每一步逻辑,才能真正理解联邦学习的本质约束与优化方向。
# 克隆仓库
git clone https://github.com/shaoxiongji/federated-learning.git
cd federated-learning
# 安装依赖(注意:torch==0.4.1 可能需要 conda 环境兼容)
pip install torch==0.4.1 torchvision==0.2.1 matplotlib numpy scikit-learn
# 运行 MNIST IID 实验
python main_fed.py --epochs 5 --num_users 10 --dataset mnist --model cnn --iid
# 运行 MNIST Non-IID 实验
python main_fed.py --epochs 5 --num_users 10 --dataset mnist --model cnn
# 运行 CIFAR10 实验
python main_fed.py --epochs 5 --num_users 10 --dataset cifar --model cnn --iid
实验结果(训练曲线、精度)自动保存到 save/ 目录。
本分析基于 GitHub 仓库代码(master分支)及项目 README 编写,图片因代理限制无法验证,未附图。