Federated-Learning-in-PyTorch
PyTorch 实现的联邦学习研究框架,支持 8 种算法、20+ 数据集与 7 种 Non-IID
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
PyTorch 实现的联邦学习研究框架,支持 8 种算法、20+ 数据集与 7 种 Non-IID
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
当你需要在来自数千部手机的数据上训练一个统一的预测模型时,传统的机器学习流程要求把所有数据汇聚到一台服务器上——但这些数据可能是你的微信聊天记录、银行的交易日志、医院的病历档案。数据不能离开设备,但模型必须共同进步,这就是联邦学习(Federated Learning)诞生的核心矛盾。
vaseline555/Federated-Learning-in-PyTorch 正是为解决这一矛盾而生的研究级实现。它用 PyTorch 将联邦学习的整个研究流程封装为模块化、可配置的实验框架,让研究者无需从零搭建基础设施,就能快速验证新算法在不同数据异构场景下的效果。
这个仓库的作者 vaseline555(Seok-Ju Hahn)在 2023 年 4 月于 Reddit 发布了这个项目的重构版本,立即在 r/MachineLearning 社区引发关注。项目最初源于个人的研究需求——作者在撰写联邦学习方向论文时,发现现有实现要么过于工程化(缺少可扩展接口),要么过于简单(缺少对真实非独立同分布场景的支持)。经过数月的重构,形成了今天这个兼顾科研严谨性与工程可用性的统一框架。
截至目前,该项目已获得 471 颗星、89 个 Fork、4 个 open issue,在 GitHub 联邦学习 PyTorch 实现中属于高影响力项目。其 topics 覆盖 FedAvg、FedProx、FedOpt 等主流算法标签,表明其学术认可度较高。
本项目的数据集支持能力是其最大的亮点之一。作者实现了一套自动下载 + 分割的管线,涵盖了联邦学习研究中最常用的 benchmark:
| 类别 | 数据集 | 说明 |
|---|---|---|
| 图像分类 | MNIST、CIFAR-10、Fashion-MNIST、TinyImageNet、CINIC-10 | torchvision 全支持 |
| 文本分类 | Reddit、Sent140、FEMNIST(手写字符+数字) | LEAF benchmark 核心 |
| 时序数据 | Shakespeare(莎士比亚文本)、GLEAM(时间序列) | 真实场景模拟 |
| 表格数据 | Adult、Heart、Cover | UCI ML Repository 经典 |
| 语音 | SpeechCommands | 支持语音任务 |
关键特性:所有数据集均可通过 --dataset 参数自动下载,无需手动准备原始文件。这对于需要频繁切换数据集的研究实验来说,极大地提升了效率。
联邦学习的核心挑战在于数据分布的非独立同分布(Non-IID)特性——不同客户端的数据分布天然不同。本项目实现了 7 种数据切分策略:
这意味着研究者可以在同一个代码框架下,对比算法在 IID 基线 vs 各种 Non-IID 场景下的表现差异,这是论文投稿时的标配实验设计。
项目内置了从 logistic 回归到 Transformer 的完整模型生态:
研究者无需修改模型代码,只需传参即可切换 backbone。
项目实现了学术界最常用的联邦学习算法套件:
| 算法 | 论文 | 核心思想 |
|---|---|---|
| FedAvg | McMahan et al., 2016 | 本地 SGD 多轮迭代 + 模型平均 |
| FedSGD | McMahan et al., 2016 | 单轮本地梯度 + 聚合 |
| FedAvgM | Hsu et al., 2019 | 引入动量缓解异质数据问题 |
| FedProx | Li et al., 2020 | 带正则项处理异构系统 |
| FedYogi | Reddi et al., 2020 | 自适应学习率(Adam 变体) |
| FedAdam | Reddi et al., 2020 | Adam 风格的联邦优化器 |
| FedAdaGrad | Reddi et al., 2020 | AdaGrad 风格的联邦优化器 |
| Fedyogi | 来自 repository | Yogi 优化器变体 |
每种算法都有对应的 Server 类(服务端聚合逻辑)和 Client 类(客户端本地训练逻辑),遵循统一的抽象接口,便于扩展新算法。
项目的代码架构遵循经典的 Client-Server 联邦学习模式,整体分为三层:
src/
├── algorithm/ # 联邦优化算法(FedAvg、FedProx 等)
├── client/ # 客户端实现(本地训练、评估、上传)
├── server/ # 服务端实现(模型分发、聚合、评估)
├── datasets/ # 数据集下载、预处理、切分
├── models/ # 20+ 预置神经网络模型
├── loaders/ # 数据加载器、模型加载器
└── metrics/ # 评估指标管理
main.py # 入口脚本,参数解析 + 实验编排
核心执行流程(main.py):
load_dataset(args):根据参数加载并切分数据集,生成 K 个客户端数据子集load_model(args):根据参数加载指定模型架构server_class(...):动态加载算法对应的 Server 类(如 FedavgServer)_sample_clients)client.update())client.upload())_aggregate),更新全局模型_central_evaluate + server.evaluate)架构亮点:采用了 ABC(Abstract Base Class)抽象基类 定义统一接口。BaseServer 和 BaseClient 分别定义了服务端和客户端的抽象方法,任何新算法的实现只需继承这两个基类,遵循接口规范即可。
项目通过 argparse 提供丰富的命令行参数,主要分为以下几类:
数据配置:--dataset、--split_type(iid/patho/dirichlet)、--alpha(Dirichlet 浓度参数)、--K(客户端数量)、--C(采样比例)
算法配置:--algorithm(fedavg/fedprox/fedsgd 等)、--E(本地 epoch 数)、--B(batch size)
训练配置:--R(联邦总轮数)、--lr、--lr_decay、--optimizer、--criterion
日志配置:--use_tb(开启 TensorBoard 实时可视化)、--log_path、--result_path
示例命令(复现 FedAvg 原始论文 MNIST 实验):
python3 main.py --exp_name "FedAvg_MNIST_CNN_Patho_C0.1_B10" --dataset MNIST --model_name TwoCNN --algorithm fedavg --split_type patho --K 100 --R 1000 --E 5 --C 0.1 --B 10 --optimizer SGD --lr 0.1 --lr_decay 0.99 --use_tb --device cuda
这行命令就可以完整复现 FedAvg 论文中 Pathological Non-IID 设置下的 MNIST 实验,并实时在 TensorBoard 中查看准确率曲线。
联邦学习虽避免了原始数据的直接传输,但模型参数本身也会泄露数据信息。该项目目前未集成任何差分隐私(Differential Privacy)或安全多方计算(MPC)机制,学术研究中若涉及敏感数据,需自行引入隐私保护层。
所有客户端-服务端通信均为明文传输,无加密压缩优化。对于大规模 K=1000 客户端的实验,模型参数传输的通信开销可能成为瓶颈(尤其在无线网络环境下)。
本项目的「联邦」是模拟分布式——所有客户端代码运行在同一台机器上,通过多线程/进程模拟不同设备。这对于验证算法正确性完全够用,但无法评估真实的网络延迟、设备异构性(算力差异、断线率)对系统的影响。如需真实分布式部署,建议使用 PySyft 或 FATE 平台。
联邦学习正在从学术研究走向产业落地:
本项目作为联邦学习研究的「实验基础设施」,在学术引用和工业参考方面都有重要价值。随着数据隐私法规(GDPR、中国数据安全法)的日趋严格,联邦学习从「nice to have」变成「must have」,这类工具库的需求会持续增长。
# 克隆仓库
git clone https://github.com/vaseline555/Federated-Learning-in-PyTorch.git
cd Federated-Learning-in-PyTorch
# 安装依赖
pip install -r requirments.txt
# 额外安装 PyTorch(按你的 CUDA 版本选择)
conda install pytorch==1.12.0 torchvision==0.13.0 torchaudio==0.12.0 cudatoolkit=11.6 -c pytorch -c conda-forge
# 运行示例(MNIST + FedAvg + CNN)
python3 main.py --exp_name test_fedavg_mnist --dataset MNIST --model_name TwoCNN --algorithm fedavg --split_type iid --K 10 --R 50 --E 5 --C 0.1 --B 10 --optimizer SGD --lr 0.1 --device cpu --use_tb
# 打开 TensorBoard 查看训练曲线
tensorboard --logdir ./log --port 6006
最低硬件需求:CPU 即可运行小规模实验(K≤10);如需 GPU 加速(推荐),需要 NVIDIA 显卡 + CUDA 环境。