新闻详情

新闻详情

首页 / 资讯中心 / 详情

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

发布时间:2026/9/28 6:26:30来源:尧图网络
联邦学习实验复现指南: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_modeltransfer.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, fracparticipation_rate) # 2. 下发全局模型各客户端本地训练 local_weights [] for client in selected_clients: w FedOur_LocalUpdate.train( global_model, client.data, lrlocal_lr, epochslocal_epochs ) local_weights.append(w) # 3. 服务端聚合 global_model FedOur_Aggr.aggregate(local_weights, weightsclient_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.01local_epochs在 Cifar-10 实验里一般设 1 到 5设大了就是血泪经验里的「客户端漂移」重灾区。2.3 三个实验的配置差异这份资源的三个实验不是简单换数据集而是各自验证不同维度的问题。实验一在 Cifar-10 上对比 FedAvg、FedPer(Classify)、FedPer(Classify 1 Block)、FedRep(Classify) 和 FedOur 五种方法。关键差异在于个性化层的设计FedPer 只个性化分类层FedPer1Block 多解冻一个残差块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 环境确认这份代码基于 PyTorchrequirements.txt里列了核心依赖。我一般会先建一个干净的虚拟环境避免和系统里的包打架。# 创建虚拟环境以 conda 为例 conda create -n fedexp python3.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( rootDATA_ROOT, trainTrue, downloadTrue, transformtrain_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 的downloadTrue会自动下载到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 上的结果图但你要自己跑一遍才能理解背后的趋势。# dermamnist10 客户端 python FedOur.py --dataset dermamnist --clients 10 --rounds 100 --algorithm fedavg # dermamnist50 客户端 python FedOur.py --dataset dermamnist --clients 50 --rounds 100 --algorithm fedavg # dermamnist100 客户端 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.110 客户端时每轮只有 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_lr0.01, inner_steps5): # 复制全局模型的基本层 base_model copy.deepcopy(global_model.base_layers) # 在目标域上做少量梯度步的元学习 optimizer torch.optim.SGD(base_model.parameters(), lrinner_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.05inner_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], labelFedAvg) plt.plot(log[rounds], log[loss], labelLoss) plt.xlabel(Communication Round) plt.ylabel(Accuracy / Loss) plt.legend() plt.savefig(my_cifar10_curve.png, dpi150)逻辑说明把每轮评估的准确率和损失存成 JSON 或 CSV再用matplotlib画图这样你可以自由调整横轴范围和样式。常见做法是每个算法存一个日志文件最后在一张图上叠加对比。注意plt.savefig的dpi设 150 以上论文或报告里才清晰。参数说明log[rounds]是轮数列表log[acc]是对应准确率。如果日志里没有存轮数可以用range(len(acc))代替。画图时plt.legend()的位置可以用loclower 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, transformNone): 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里训练集用RandomCropRandomHorizontalFlipNormalize测试集只用ResizeNormalize。Normalize的均值和方差用你自己数据集的统计值不要直接套 Cifar-10 的。最后验证迁移效果。先在小样本上跑 10 轮确认 loss 在降、准确率在升再放大到全量和 100 轮。如果 10 轮内 loss 完全不降大概率是学习率或数据预处理有问题别急着加轮数。从那以后我每次迁移新数据集都强制先跑一个 10 轮的小实验确认流程通畅再投入长时间训练。希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

不着急管住嘴多喝水迈开腿洗干净睡好觉:六个健康动作的系统拆解 2026/9/28 7:17:49

不着急管住嘴多喝水迈开腿洗干净睡好觉:六个健康动作的系统拆解

“不着急、管住嘴、多喝水、迈开腿、洗干净、睡好觉”,这句话我第一次是在小区门口的健康宣传栏看到的,当时扫了一眼没当回事,觉得就是句给老年人听的顺口溜。直到前阵子体检报告出了一串箭头,我才把它翻出来认真琢磨,…

阅读更多 →
class-transformer 基础用法指南:plainToInstance 与 instanceToPlain 核心转换函数与装饰器详解 2026/9/28 7:17:48

class-transformer 基础用法指南:plainToInstance 与 instanceToPlain 核心转换函数与装饰器详解

序列化后端前端 【免费下载链接】class-transformer Decorator-based transformation, serialization, and deserialization between objects and classes. 项目地址: https://gitcode.com/gh_mirrors/cl/class-transformer 点击查看 免费下载 本篇指南围绕 class…

阅读更多 →
国外做gif的网站新手入门:3步搞定服务器与域名 2026/9/28 7:17:35

国外做gif的网站新手入门:3步搞定服务器与域名

国外做gif的网站新手入门:3步搞定服务器与域名 别被“域名服务器搞不懂”这四个字劝退,这才是新手入门建站最大的拦路虎。很多想在国外做gif的网站的朋友,卡在第一步注册VPS和解析DNS上就放弃了。其实逻辑很简单:域名是门牌号,服务器是房子…

阅读更多 →
PLC ST语言定时器实战:TON/TOF指令原理与工程应用 2026/9/28 7:17:29

PLC ST语言定时器实战:TON/TOF指令原理与工程应用

做PLC项目调试,最头疼的往往不是逻辑本身多复杂,而是设备动作的时序对不上。拿ST语言写定时器控制,稍微有一点经验的人都绕不开TON和TOF这两个指令。TON是接通延时定时器,IN端有信号了并不马上输出,而是等计时到设定值…

阅读更多 →
ST语言定时器全解析:TON/TOF原理、应用与排错技巧 2026/9/28 7:17:29

ST语言定时器全解析:TON/TOF原理、应用与排错技巧

做PLC项目的人应该都有同感:梯形图里最常用的指令,除了常开常闭触点,就是定时器。我刚从梯形图转ST语言那会儿,最别扭的就是定时器——梯形图里拖一个TON框出来,填个时间就完事;换成ST之后,不少…

阅读更多 →
基于Python+Hadoop的气象分析大屏可视化毕设全流程指南 2026/9/28 7:17:29

基于Python+Hadoop的气象分析大屏可视化毕设全流程指南

上个答辩季,我帮好几个学弟学妹远程排过这类“基于PythonHadoop的气象分析大屏可视化”项目的坑。说实话,这个题目在近年来算是大数据方向毕业设计里相当能打的一种组合:既有Hadoop生态的重量感,又有大屏可视化带来的直接观感冲击…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉