news 2026/9/28 6:26:23

联邦学习实验复现指南:FedAvg到FedOur三组对比实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
联邦学习实验复现指南:FedAvg到FedOur三组对比实战

简介:本资源是一套基于Python实现的联邦学习实验项目,面向人工智能、计算机及相关专业的学生、教师与企业员工,适合作为毕设、课程设计或算法入门进阶的实战参考。项目围绕FedAvg、FedPer、FedRep与FedOur等算法展开三个实验:在Cifar-10上对比各方法的准确率与目标损失,在MedMNIST上测试10、50、100等不同客户端数量下的表现,并在Chest X-Ray Images数据集上验证全局模型与本地模型经Meta-Transfer训练的效果。压缩包共43个文件,包含14个Python源码、18张png与2张jpg实验曲线图、5个xml配置及说明文档,整体约631KB,目录涵盖模型定义、数据采样、聚合与本地更新等模块,结构清晰。目前已有227人学习。读者可获取完整可运行代码、预置模型与可视化结果,便于复现实验、理解联邦学习流程并在此基础上二次开发。

1. 联邦学习实验复现:从 FedAvg 到 FedOur 的三组对比怎么跑

如果你正在做联邦学习方向的毕设或课程设计,大概率会遇到一个尴尬局面:论文里的 FedAvg、FedPer、FedRep 公式都看得懂,但真要自己从零搭一套能跑通、能出准确率曲线、还能横向对比多个算法的实验框架,光是数据划分和客户端采样就能卡上一周。这份基于 Python 的联邦学习实验资源,核心价值就在于它把三组完整实验、可运行的源码、预训练模型和训练曲线图打包在了一起。它覆盖 Cifar-10 上的多算法对比、MedMNIST 上的客户端数量敏感性测试,以及 Chest X-Ray 上的 Meta-Transfer 微调验证。适合已经懂 PyTorch 基础、想快速拿到一套可复现联邦学习实验骨架的在校学生和初级算法工程师。下面我按实际拆包运行的顺序,把这份资源讲透。

2. 实验框架拆解:FedOur 主流程与模块依赖关系

2.1 目录结构与核心文件职责

拿到federal-learning-experiment-master.zip解压后,根目录下是典型的 Python 实验工程布局。先别急着跑,把每个文件的职责搞清楚,后面调参和排错才不会抓瞎。

文件/目录职责
FedOur_LocalUpdate.py客户端本地训练入口,负责在本地数据上执行 SGD 更新
FedOur_Aggr.py服务端聚合逻辑,实现全局模型参数加权平均
FedOur.py主控脚本,串联客户端采样、本地更新、聚合、评估全流程
dataset.py数据集加载与预处理,含 Cifar-10、MedMNIST、Chest X-Ray 的读取逻辑
sampling.py客户端采样策略,控制每轮参与训练的客户端子集
options.py全局超参数配置,学习率、轮数、客户端数都在这里改
models/模型定义,含 Resnet18、Resnet34、Nets、global_model
transfer.pyMeta-Transfer 微调逻辑,对应实验三
test.py独立评估脚本,加载模型权重跑测试集
utils/工具函数,含日志、指标计算、模型保存
img/训练曲线图,覆盖三个实验的 loss 和 acc

这个结构的好处是职责分离清晰:想改聚合策略只动FedOur_Aggr.py,想换数据集只动dataset.py,想调超参只动options.py。常见做法是先把options.py通读一遍,把所有默认参数记下来,再决定改哪些。

2.2 联邦学习主循环的数据流

联邦学习的核心循环可以用一句话概括:采样客户端 → 下发全局模型 → 本地训练 → 上传更新 → 聚合 → 评估。这份代码里,FedOur.py是主控,它依次调用sampling.py选客户端、FedOur_LocalUpdate.py做本地训练、FedOur_Aggr.py做聚合。

# FedOur.py 主循环伪代码(基于项目结构还原) for round in range(global_rounds): # 1. 采样本轮参与的客户端 selected_clients = sampling.sample_clients(clients, frac=participation_rate) # 2. 下发全局模型,各客户端本地训练 local_weights = [] for client in selected_clients: w = FedOur_LocalUpdate.train( global_model, client.data, lr=local_lr, epochs=local_epochs ) local_weights.append(w) # 3. 服务端聚合 global_model = FedOur_Aggr.aggregate(local_weights, weights=client_sizes) # 4. 每轮评估 acc, loss = test.evaluate(global_model, test_loader)

逻辑说明:sampling.sample_clients控制每轮参与率,participation_rate设太小会导致训练不稳定,设太大则失去联邦学习「部分参与」的意义。FedOur_Aggr.aggregate里的weights参数是按客户端样本量加权,样本多的客户端对全局模型影响更大,这是 FedAvg 的标准做法。local_epochs是本地训练轮数,这个参数直接决定客户端漂移程度——本地训太多轮,各客户端模型会跑偏,聚合时反而拉低全局效果。

参数说明:global_rounds建议从 50 起步观察收敛趋势;local_lr通常比集中式训练小一个量级,常见取 0.01;local_epochs在 Cifar-10 实验里一般设 1 到 5,设大了就是血泪经验里的「客户端漂移」重灾区。

2.3 三个实验的配置差异

这份资源的三个实验不是简单换数据集,而是各自验证不同维度的问题。

实验一在 Cifar-10 上对比 FedAvg、FedPer(Classify)、FedPer(Classify + 1 Block)、FedRep(Classify) 和 FedOur 五种方法。关键差异在于个性化层的设计:FedPer 只个性化分类层,FedPer+1Block 多解冻一个残差块,FedRep 则把表示层和分类头分开处理。跑这个实验时,options.py里要确认algorithm字段切换正确,否则你会看到两条一模一样的曲线还以为是玄学。

实验二用 MedMNIST 的 dermamnist 和 bloodmnist 子集,测试客户端数量从 10、50 到 100 的影响。img/目录下的dermamnist_10clients_acc.png、dermamnist_50clients_acc.png、dermamnist_100clients_acc.png就是这组对比的结果。客户端越多,每轮聚合的方差越大,收敛越慢,但最终精度上限可能更高。

实验三用 Chest X-Ray 数据集,对比 FedAvg 全局模型、Local 训练本地模型,以及全局基本层经过 Meta-Transfer 微调后的效果。transfer.py是这组实验的核心,它加载全局模型的基本层,在目标域上做少量梯度步的元学习适配。

3. 环境搭建与第一个实验跑通:Cifar-10 五算法对比

3.1 依赖安装与 Python 环境确认

这份代码基于 PyTorch,requirements.txt里列了核心依赖。我一般会先建一个干净的虚拟环境,避免和系统里的包打架。

# 创建虚拟环境(以 conda 为例) conda create -n fedexp python=3.8 -y conda activate fedexp # 安装依赖 pip install -r requirements.txt # 确认 PyTorch 和 CUDA 可用 python -c "import torch; print(torch.__version__, torch.cuda.is_available())"

逻辑说明:Python 3.8 是这类毕设代码最常见的版本,太新的 Python 可能遇到部分依赖不兼容。torch.cuda.is_available()返回True才说明 GPU 可用,如果返回False,后面训练会慢到让你怀疑人生。常见做法是先在options.py里把device参数确认为cuda,没有 GPU 就改cpu,但 Cifar-10 五算法对比在 CPU 上跑完整流程可能要几个小时。

参数说明:requirements.txt里通常包含torch、torchvision、numpy、matplotlib、scikit-learn等。如果安装时遇到版本冲突,优先保证torch和torchvision版本匹配,这两个不匹配会直接报错。

3.2 数据集准备与路径配置

Cifar-10 可以通过torchvision.datasets自动下载,但 MedMNIST 和 Chest X-Ray 需要手动准备。dataset.py里一般会有数据根目录的配置项。

# dataset.py 中常见的数据路径配置 DATA_ROOT = './data' # 数据根目录,按实际存放位置修改 # Cifar-10 自动下载 train_set = torchvision.datasets.CIFAR10( root=DATA_ROOT, train=True, download=True, transform=train_transform ) # MedMNIST 需要先安装 medmnist 包 # pip install medmnist import medmnist from medmnist import INFO data_flag = 'dermamnist' info = INFO[data_flag] DataClass = getattr(medmnist, info['python_class'])

逻辑说明:Cifar-10 的download=True会自动下载到DATA_ROOT,国内网络可能较慢,可以提前手动下载好放到对应目录。MedMNIST 通过medmnist包加载,data_flag切换dermamnist或bloodmnist对应实验二的不同子集。Chest X-Ray 数据集需要自己下载后按目录结构放好,dataset.py里通常用ImageFolder读取。

参数说明:DATA_ROOT建议设为绝对路径,相对路径在不同工作目录下运行容易翻车。MedMNIST 的size参数可选 28、64、128,实验里一般用 28 或 64,设太大显存吃不消。

3.3 启动实验一:五算法对比

配置确认后,直接跑主脚本。实验一的核心是切换algorithm参数,分别跑五种方法。

# 跑 FedAvg python FedOur.py --algorithm fedavg --dataset cifar10 --clients 10 --rounds 100 # 跑 FedPer(Classify) python FedOur.py --algorithm fedper --dataset cifar10 --clients 10 --rounds 100 # 跑 FedPer(Classify + 1 Block) python FedOur.py --algorithm fedper_1block --dataset cifar10 --clients 10 --rounds 100 # 跑 FedRep(Classify) python FedOur.py --algorithm fedrep --dataset cifar10 --clients 10 --rounds 100 # 跑 FedOur python FedOur.py --algorithm fedour --dataset cifar10 --clients 10 --rounds 100

逻辑说明:每次运行会生成对应的准确率和损失日志,img/目录下的cifar-10-acc.png和cifar-10-loss.png就是这些结果的汇总图。如果你想自己复现曲线,需要把五次运行的日志分别保存,再用matplotlib画图。常见做法是加一个--log_dir参数把每次运行的结果存到不同目录,避免覆盖。

参数说明:--clients 10表示总客户端数为 10,--rounds 100是全局通信轮数。如果显存不够,把--batch_size从默认的 64 降到 32 或 16。--local_epochs控制本地训练轮数,实验一里一般设 1 到 3,设大了 FedPer 和 FedRep 的个性化优势会被削弱。

4. 实验二与实验三:客户端数量敏感性与 Meta-Transfer 微调

4.1 实验二:MedMNIST 客户端数量对比

实验二的核心变量是客户端数量。img/目录下已经给出了 10、50、100 三种客户端数在 dermamnist 和 bloodmnist 上的结果图,但你要自己跑一遍才能理解背后的趋势。

# dermamnist,10 客户端 python FedOur.py --dataset dermamnist --clients 10 --rounds 100 --algorithm fedavg # dermamnist,50 客户端 python FedOur.py --dataset dermamnist --clients 50 --rounds 100 --algorithm fedavg # dermamnist,100 客户端 python FedOur.py --dataset dermamnist --clients 100 --rounds 100 --algorithm fedavg # bloodmnist 同理,替换 --dataset 即可 python FedOur.py --dataset bloodmnist --clients 10 --rounds 100 --algorithm fedavg

逻辑说明:客户端数量增加时,每轮参与训练的客户端子集也在变。如果participation_rate固定为 0.1,10 客户端时每轮只有 1 个客户端参与,100 客户端时有 10 个。参与客户端越多,聚合后的全局模型越稳定,但单个客户端的本地数据被「稀释」的程度也越高。dermamnist_100clients_acc.png和dermamnist_10clients_acc.png的对比能直观看到这个趋势。

参数说明:--clients改的是总客户端数,sampling.py里的participation_rate控制每轮参与比例。如果显存或时间有限,可以把--rounds降到 50,观察前 50 轮的收敛趋势也够写实验报告了。MedMNIST 的图像尺寸小,--batch_size可以设大一些,64 或 128 都行。

4.2 实验三:Chest X-Ray 上的 Meta-Transfer 微调

实验三是这份资源里最有技术含量的部分。它对比三种设置:FedAvg 全局模型直接测试、Local 训练本地模型、全局基本层经过 Meta-Transfer 微调。transfer.py实现了元学习适配逻辑。

# transfer.py 中 Meta-Transfer 的核心逻辑(基于项目结构还原) def meta_transfer(global_model, target_loader, inner_lr=0.01, inner_steps=5): # 复制全局模型的基本层 base_model = copy.deepcopy(global_model.base_layers) # 在目标域上做少量梯度步的元学习 optimizer = torch.optim.SGD(base_model.parameters(), lr=inner_lr) for step in range(inner_steps): for x, y in target_loader: logits = base_model(x) loss = F.cross_entropy(logits, y) optimizer.zero_grad() loss.backward() optimizer.step() # 返回微调后的模型 return base_model

逻辑说明:Meta-Transfer 的思路是保留全局模型的基本层(特征提取部分),只在目标域上做少量梯度更新,让模型快速适配新分布。inner_lr是内循环学习率,inner_steps是内循环步数。这两个参数直接决定适配程度:步数太少欠拟合,步数太多会过拟合到目标域的小样本上。fine-tune-test-acc.jpg和fine-tune-test-loss.jpg就是这组实验的结果。

参数说明:inner_lr常见取 0.01 到 0.05,inner_steps取 3 到 10。Chest X-Ray 数据集类别不平衡,评估时除了准确率还要看每类召回率,test.py里如果有classification_report输出就更方便。

4.3 结果复现与曲线绘制

三个实验跑完后,你需要把日志整理成曲线图。img/目录下的图是作者跑出来的参考结果,你自己跑的结果应该趋势一致但具体数值会有差异。

# 绘制准确率曲线的常见做法 import matplotlib.pyplot as plt import json # 假设日志存为 json,含每轮的 acc 列表 with open('logs/fedavg_cifar10.json') as f: log = json.load(f) plt.plot(log['rounds'], log['acc'], label='FedAvg') plt.plot(log['rounds'], log['loss'], label='Loss') plt.xlabel('Communication Round') plt.ylabel('Accuracy / Loss') plt.legend() plt.savefig('my_cifar10_curve.png', dpi=150)

逻辑说明:把每轮评估的准确率和损失存成 JSON 或 CSV,再用matplotlib画图,这样你可以自由调整横轴范围和样式。常见做法是每个算法存一个日志文件,最后在一张图上叠加对比。注意plt.savefig的dpi设 150 以上,论文或报告里才清晰。

参数说明:log['rounds']是轮数列表,log['acc']是对应准确率。如果日志里没有存轮数,可以用range(len(acc))代替。画图时plt.legend()的位置可以用loc='lower right'调整,避免挡住曲线。

5. 避坑与排查:联邦学习实验里最容易翻车的五个点

5.1 客户端采样后数据为空

现象:运行时报ZeroDivisionError或ValueError: empty tensor,堆栈指向FedOur_Aggr.py的聚合函数。

原因:sampling.py里按比例采样时,如果participation_rate设得太小,或者客户端总数太少,某轮可能采到 0 个客户端。另一个常见原因是客户端数据划分时某个客户端分到了空数据集。

解决:在sampling.py里加一个保护,确保每轮至少采到 1 个客户端。数据划分时检查每个客户端的样本数,小于batch_size的客户端要么合并要么剔除。

5.2 本地训练 loss 不降反升

现象:客户端本地训练的 loss 在几个 epoch 后开始上升,全局模型准确率震荡不收敛。

原因:本地学习率local_lr设太大,或者local_epochs设太多导致客户端漂移。联邦学习里每个客户端只看到局部数据,本地训太多轮会让模型过度拟合本地分布,聚合时反而互相抵消。

解决:把local_lr降到 0.01 或更低,local_epochs控制在 1 到 3。如果还是震荡,检查数据是否需要归一化,Cifar-10 和 MedMNIST 的预处理方式不同,dataset.py里的transform要对应修改。

5.3 MedMNIST 加载报错或标签维度不对

现象:medmnist包导入失败,或者标签 shape 是(N, 1)导致cross_entropy报错。

原因:medmnist包的版本差异,旧版返回的标签是二维的,新版可能已经修复。另外INFO[data_flag]的 key 拼写错误也会导致加载失败。

解决:先pip install medmnist --upgrade升级到最新版。标签维度问题在dataset.py里加一句labels = labels.squeeze().long()即可。data_flag确认拼写为dermamnist、bloodmnist等官方名称。

5.4 GPU 显存不足导致训练中断

现象:RuntimeError: CUDA out of memory,训练在某个 batch 突然崩掉。

原因:batch_size太大,或者模型(Resnet34)比 Resnet18 更吃显存。实验三的 Chest X-Ray 图像尺寸可能比 Cifar-10 大,显存占用更高。

解决:把--batch_size减半,或者换用 Resnet18。如果还不行,在options.py里加torch.cuda.empty_cache(),或者用--device cpu先跑通流程再换 GPU。

5.5 聚合后模型参数形状不匹配

现象:FedOur_Aggr.py里state_dict加载时报size mismatch。

原因:不同客户端的模型结构不一致,比如 FedPer 和 FedAvg 的模型层数不同,聚合时直接按 key 平均会出错。另一个原因是模型保存和加载时 key 的前缀不一致(比如多了module.)。

解决:聚合前先检查所有客户端的state_dictkey 是否一致。如果前缀不一致,用collections.OrderedDict重命名。FedPer 这类有个性化层的算法,聚合时只聚合共享层,个性化层保留在本地。

6. 进阶技巧:把 FedOur 迁移到自己的数据集上

跑通三个实验只是第一步,真正有价值的是把这套框架迁移到自己的数据上。我一般会按下面的顺序操作,避免一上来就改得面目全非。

先确认数据格式。dataset.py里每个数据集对应一个load_xxx函数,你的数据如果是图像分类,按ImageFolder的目录结构放好最省事:train/class_name/image.jpg。如果是其他格式,仿照dataset.py里 Cifar-10 的写法,实现一个返回(image, label)的Dataset子类。

然后改options.py里的数据集名称和类别数。类别数错了会在模型最后一层报维度错误,这个坑很隐蔽,因为报错信息不会直接告诉你「类别数不对」。常见做法是先在dataset.py里打印一下len(dataset.classes),确认和options.py里的num_classes一致。

接着调客户端划分策略。sampling.py里默认可能是 IID 划分,如果你的数据天然按用户分组(比如每个医院一个客户端),直接把分组逻辑替换进去。非 IID 划分下,local_epochs要适当降低,否则客户端漂移会更严重。

# 自定义数据集加载的骨架 class MyDataset(Dataset): def __init__(self, root, transform=None): self.samples = [] # 存 (path, label) self.transform = transform # 遍历目录填充 samples for label, cls in enumerate(sorted(os.listdir(root))): cls_dir = os.path.join(root, cls) for img_name in os.listdir(cls_dir): self.samples.append((os.path.join(cls_dir, img_name), label)) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label = self.samples[idx] img = Image.open(path).convert('RGB') if self.transform: img = self.transform(img) return img, label

逻辑说明:这个骨架兼容ImageFolder的目录结构,sorted(os.listdir(root))保证类别顺序稳定,否则每次运行标签映射可能变,导致结果不可复现。convert('RGB')处理灰度图或 RGBA 图,避免通道数不匹配。

参数说明:transform里训练集用RandomCrop+RandomHorizontalFlip+Normalize,测试集只用Resize+Normalize。Normalize的均值和方差用你自己数据集的统计值,不要直接套 Cifar-10 的。

最后验证迁移效果。先在小样本上跑 10 轮,确认 loss 在降、准确率在升,再放大到全量和 100 轮。如果 10 轮内 loss 完全不降,大概率是学习率或数据预处理有问题,别急着加轮数。从那以后我每次迁移新数据集,都强制先跑一个 10 轮的小实验确认流程通畅,再投入长时间训练。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/28 6:25:50

矩阵系统好用才是硬道理:选型避坑与实操指南

矩阵系统这个圈子,最近确实有点热闹过头了。只要打开任何一个运营类社群,准能看见有人在推某某矩阵系统、某某多账号管理工具,搞得好像不上一套系统就没法做运营了一样。但以我这些年用过的、拆解过的、甚至帮人擦过屁股的矩阵软件来看&#…

作者头像 李华
网站建设 2026/9/28 6:25:48

一键切换 DLSS 版本:DLSS Swapper 完整安装指南

一键切换 DLSS 版本:DLSS Swapper 完整安装指南 【免费下载链接】dlss-swapper 项目地址: https://gitcode.com/GitHub_Trending/dl/dlss-swapper 你的游戏 DLSS 版本卡在 2.1.39.0,画面偶发闪烁,官方补丁却迟迟没有消息。DLSS Swapp…

作者头像 李华
网站建设 2026/9/28 6:25:31

FPGA以太网开发:Tri Mode Ethernet MAC与AXI Ethernet Subsystem选型指南

在FPGA以太网开发这条路上,选IP核这件事看起来不起眼,实际上能决定你后面三个月是顺风顺水还是天天抓bug。我见过太多项目,板子画好了、PHY选好了、时钟也规划完了,结果卡在IP核选型上——有人用Tri Mode Ethernet MAC搭好了千兆链…

作者头像 李华
网站建设 2026/9/28 6:24:09

ROS+PX4+Gazebo无人机仿真深度调优指南

1. 为什么“ROSPX4Gazebo”组合至今仍是无人机仿真不可绕过的铁三角?你刚在Ubuntu 22.04上敲完sudo apt install ros-humble-desktop,终端回显“Done”,心里一松——ROS装好了。可当你打开QGroundControl,加载PX4固件,…

作者头像 李华