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

发布时间:2026/10/7 8:45:06
联邦学习实验复现指南:FedAvg到FedOur三组对比与避坑 简介本资源是一套基于Python实现的联邦学习实验项目面向人工智能、计算机及相关专业的在校学生、教师与企业员工适合作为毕业设计、课程设计或项目立项演示也可供进阶学习者参考。项目围绕Cifar-10、MedMNIST与Chest X-Ray Images三个数据集展开对比FedAvg、FedPer、FedRep与FedOur等算法在准确率和目标损失上的表现并测试不同客户端数量下的训练效果同时涉及全局模型与本地模型的Meta-Transfer思路。压缩包共43个文件包含14个Python源码文件、18张png与2张jpg实验曲线图、5个xml配置及md说明文档等整体约631KB目录涵盖模型定义、数据采样、训练与聚合脚本等模块。目前已有227人学习下载。读者可获得完整可运行的联邦学习实验代码、模型文件与结果图示便于复现实验、理解算法差异并在此基础上修改扩展。1. 联邦学习实验复现从 FedAvg 到 FedOur 的三组对比怎么跑起来联邦学习这个词这两年热度一直不低但真正把代码拉下来能跑通、能复现出论文级对比曲线的人并不多。这份基于 Python 的联邦学习实验包核心价值就在于它把三组实验的完整流程都固化成了可执行脚本Cifar-10 上 FedAvg、FedPer、FedRep 与 FedOur 的准确率和损失对比MedMNIST 上客户端数量从 10 到 100 的扩展性测试以及 Chest X-Ray 上全局模型与本地模型经 Meta-Transfer 微调后的效果差异。它适合正在做毕设、课程设计或者想快速验证联邦学习算法改动的同学和工程师。整个包不依赖复杂框架纯 PyTorch 加几个自定义模块跑起来门槛不高但里面关于数据划分、客户端采样和模型聚合的细节恰恰是决定实验能不能复现的关键。2. 环境搭建与数据准备把依赖和数据管道先理顺2.1 依赖安装与目录结构确认拿到压缩包后先别急着跑脚本第一步是把环境对齐。包里带了requirements.txt但联邦学习实验对 PyTorch 版本和 CUDA 匹配比较敏感我一般会先建一个干净的虚拟环境再按文件装依赖。# 创建虚拟环境Python 3.8 是这类实验比较稳的版本 python -m venv fed_env source fed_env/bin/activate # Windows 下用 fed_env\Scripts\activate # 安装依赖建议先升级 pip 避免解析冲突 pip install --upgrade pip pip install -r requirements.txt # 确认 PyTorch 和 torchvision 能正常调用 GPU python -c import torch; print(torch.__version__, torch.cuda.is_available())这里有个参数要留意requirements.txt里如果写的是torch1.8这种宽松约束实际装出来的可能是最新版而新版 PyTorch 在某些旧版 torchvision 搭配下会报算子不兼容。稳妥做法是手动指定一个组合比如torch1.12.1配torchvision0.13.1CUDA 版本按你显卡驱动选 cu113 或 cu116。装完后检查torch.cuda.is_available()返回 True否则后面训练会默认落到 CPU一个 epoch 能跑十几分钟直接拖垮实验节奏。目录结构上核心文件集中在根目录和models/下。FedOur.py是主入口FedOur_LocalUpdate.py负责客户端本地训练FedOur_Aggr.py做服务器端聚合dataset.py和sampling.py管数据加载与客户端采样options.py集中了所有超参数。models/里放了Resnet18.py、Resnet34.py、Nets.py和global_model.py分别对应不同骨干网络和全局模型定义。img/下那批 png 和 jpg 是作者跑出来的损失和准确率曲线可以当作你复现时的对照基准。2.2 数据集下载与目录摆放三组实验用了三个数据集Cifar-10、MedMNIST 和 Chest X-Ray Images。Cifar-10 通常由 torchvision 自动下载但 MedMNIST 和 Chest X-Ray 需要手动准备。# dataset.py 里通常会有类似这样的加载逻辑确认路径指向你实际存放数据的位置 import torchvision.datasets as datasets # Cifar-10 自动下载到指定 root cifar_train datasets.CIFAR10(root./data/cifar10, trainTrue, downloadTrue) cifar_test datasets.CIFAR10(root./data/cifar10, trainFalse, downloadTrue) # MedMNIST 一般通过 medmnist 库加载注意指定 subset 和 size import medmnist from medmnist import INFO data_flag dermamnist # 实验里用到了 dermamnist 和 bloodmnist info INFO[data_flag] DataClass getattr(medmnist, info[python_class]) train_dataset DataClass(splittrain, downloadTrue, root./data/medmnist)MedMNIST 的下载有时候会因为网络问题卡住常见做法是提前用浏览器或下载工具把 npz 文件拉到./data/medmnist下再让代码走本地加载。Chest X-Ray Images 数据集在实验三里用到需要按类别分文件夹存放dataset.py里一般用ImageFolder读取目录结构得是train/NORMAL/、train/PNEUMONIA/这种形式否则标签会对不上。提示数据路径最好在options.py里统一改不要散落在各个脚本里硬编码不然后面换数据集要满项目找路径。2.3 客户端划分与采样参数联邦学习实验里客户端怎么划分、每轮采样多少客户端直接决定对比曲线是否公平。sampling.py负责这部分逻辑常见做法是 IID 或 Dirichlet 非独立同分布划分。# sampling.py 中典型的客户端采样逻辑 import numpy as np def sample_clients(num_clients, fraction, seed42): 每轮按比例采样客户端 num_clients: 总客户端数实验二里会设成 10、50、100 fraction: 每轮参与比例通常 0.1 到 0.5 rng np.random.default_rng(seed) num_selected max(int(num_clients * fraction), 1) selected rng.choice(num_clients, num_selected, replaceFalse) return selectedfraction这个参数在实验二里特别关键。客户端数从 10 涨到 100 时如果 fraction 固定为 0.1那每轮参与数就从 1 变成 10聚合行为会明显不同。作者在img/里放了dermamnist_10clients_acc.png、dermamnist_50clients_acc.png、dermamnist_100clients_acc.png三组曲线复现时建议先固定 fraction 再改客户端数否则准确率波动里混了采样比例的影响归因会乱。3. 三组实验的运行方式与参数配置3.1 实验一Cifar-10 上五种算法对比实验一的目标是把 FedAvg、FedPer(Classify)、FedPer(Classify 1 Block)、FedRep(Classify) 和 FedOur 放在同一套数据划分下比准确率和目标损失。入口在FedOur.py通过options.py里的参数切换算法。# 跑 FedAvg 基线 python FedOur.py --dataset cifar10 --algorithm fedavg --num_clients 100 --fraction 0.1 --epochs 100 --lr 0.01 # 跑 FedPer注意 personalize 层配置 python FedOur.py --dataset cifar10 --algorithm fedper --personalize classify --num_clients 100 --fraction 0.1 # 跑 FedOur作者自己的方法 python FedOur.py --dataset cifar10 --algorithm fedour --num_clients 100 --fraction 0.1 --meta_lr 0.001--algorithm控制走哪条分支--personalize决定个性化层是只留分类头还是加一个残差块。FedPer(Classify 1 Block) 和 FedPer(Classify) 的差别就在这个参数上前者多保留一个 block 做本地个性化后者只留分类层。--meta_lr是 FedOur 里元学习部分的步长设太大容易震荡设太小收敛慢作者在img/cifar-10-loss.png和cifar-10-acc.png里的曲线对应的是 0.001 这个量级。跑的时候注意看FedOur_LocalUpdate.py里的本地 epoch 数。联邦学习里本地训练轮数local_epoch是个双刃剑设太小客户端欠拟合全局模型学不到东西设太大客户端漂移严重聚合后反而掉点。Cifar-10 上常见做法是local_epoch5配合lr0.01和 SGD 优化器。3.2 实验二MedMNIST 客户端数量扩展性测试实验二换到 MedMNIST 的 dermamnist 和 bloodmnist 两个子集重点看客户端数量变化对收敛的影响。这个实验的脚本调用方式和实验一类似但数据集参数和客户端数要改。# dermamnist10 个客户端 python FedOur.py --dataset dermamnist --num_clients 10 --fraction 0.5 --epochs 50 --algorithm fedour # dermamnist50 个客户端 python FedOur.py --dataset dermamnist --num_clients 50 --fraction 0.1 --epochs 50 --algorithm fedour # bloodmnist100 个客户端 python FedOur.py --dataset bloodmnist --num_clients 100 --fraction 0.1 --epochs 50 --algorithm fedour这里有个容易翻车的点客户端数变了每轮参与数也跟着变如果不同实验之间想公平对比要么固定每轮参与客户端数要么在论文里明确说明采样策略。作者在img/里放了dermamnist_10clients_loss.png到dermamnist_100clients_acc.png一整套曲线复现时建议先把 10 客户端的结果跑出来和图片对一下趋势再往上加客户端数。MedMNIST 图像分辨率低ResNet18 足够用换 ResNet34 反而容易过拟合models/Resnet18.py和Resnet34.py都在但实验二默认走 18。3.3 实验三Chest X-Ray 上的 Meta-Transfer 微调实验三比较的是 FedAvg 全局模型、Local 训练本地模型以及全局基本层经过 Meta-Transfer 后的效果。这部分逻辑在transfer.py和FedOur.py的微调分支里。# 先跑 FedAvg 得到全局模型 python FedOur.py --dataset chestxray --algorithm fedavg --num_clients 50 --epochs 30 --save_global # 再做 Meta-Transfer 微调 python transfer.py --dataset chestxray --global_ckpt ./checkpoints/fedavg_global.pth --meta_epochs 20 --inner_lr 0.01 --outer_lr 0.001--inner_lr和--outer_lr是元学习的内外循环步长transfer.py里一般会先复制全局模型参数再在支持集上做几步内更新最后在查询集上算元梯度。img/fine-tune-test-loss.jpg和fine-tune-test-acc.jpg就是这组实验的曲线。Chest X-Ray 数据量比 MedMNIST 大显存不够的话把 batch size 降到 16 或 8同时把--num_clients调小不然每轮聚合等待时间会很长。注意实验三的 Meta-Transfer 部分对随机种子敏感options.py里如果有--seed参数复现时固定成同一个值否则曲线抖动会让你怀疑代码有问题。4. 避坑与排查复现联邦学习实验时最容易翻车的五件事4.1 损失曲线不下降准确率卡在随机水平现象跑了几十个 epochloss一直在 2.3 附近晃准确率 10% 左右和 Cifar-10 随机猜差不多。原因最常见的是数据标签没对齐。dataset.py里如果用了ImageFolder但目录名和标签映射写错或者 MedMNIST 的split参数写成了train却拿去做测试标签就全乱了。另一个可能是客户端采样后本地数据为空sampling.py里num_selected算出来是 0本地训练直接跳过。解决先在dataset.py里加一行打印确认len(train_dataset)和train_dataset[0][1]的标签范围。再检查sampling.py的fraction是否过小导致int(num_clients * fraction)为 0把max(..., 1)加上。4.2 聚合时张量形状不匹配报错现象FedOur_Aggr.py里做参数平均时抛RuntimeError: The size of tensor a must match the size of tensor b。原因不同客户端的模型结构不一致。FedPer 和 FedOur 里有个性化层如果某些客户端保留了额外 block某些没保留聚合时 state_dict 的 key 就对不上。另外models/global_model.py里定义的全局模型和Resnet18.py里的本地模型如果层名不同加载也会出问题。解决聚合前先打印每个客户端 state_dict 的 keys确认结构一致。个性化层在聚合时要排除只聚合全局部分。FedOur_Aggr.py里通常有个exclude_keys列表把personalize相关的 key 加进去。4.3 GPU 显存溢出训练中途崩掉现象跑实验二 100 客户端时几个 epoch 后报CUDA out of memory。原因客户端数多如果每轮把所有客户端模型都留在显存里占用会线性增长。另外 MedMNIST 虽然图小但ResNet34比ResNet18参数量翻倍显存不够时优先换回 18。解决在FedOur_LocalUpdate.py里确保每个客户端训练完就释放中间变量用del加torch.cuda.empty_cache()。batch size 从 64 降到 32 或 16。如果还不够把--num_clients先降到 50 跑通再往上加。4.4 复现曲线和 img 里的对不上现象自己跑出来的准确率比img/cifar-10-acc.png里低好几个点。原因随机种子没固定、学习率调度不一致、或者数据增强策略不同。联邦学习里客户端采样本身有随机性不同 seed 下曲线抖动几个点是正常的但如果差 10 个点以上多半是超参没对齐。解决在options.py里找--seed固定成作者用的值常见是 42 或 0。检查--lr和--local_epoch是否和 README 里写的一致。数据增强部分Cifar-10 上常见的是 RandomCrop 加 RandomHorizontalFlip如果作者代码里关了增强你开了结果也会偏。4.5 MedMNIST 下载失败或加载报错现象medmnist库下载 npz 时卡住或者加载时报KeyError: dermamnist。原因网络问题导致下载不完整或者 medmnist 版本和数据集名称不匹配。旧版 medmnist 里 dermamnist 的 key 可能不一样。解决手动下载 npz 放到./data/medmnist下代码里把downloadTrue改成False。确认medmnist.__version__必要时pip install medmnist2.2.1这种指定版本。INFO字典的 key 打印出来看一眼确认dermamnist和bloodmnist都在。5. 进阶用法把 FedOur 接到自己的模型上并验证聚合正确性跑通三组实验只是第一步这份代码更大的价值在于它把联邦学习的本地更新、聚合、元学习微调拆成了独立模块你可以把自己的模型塞进去。我一般会按这个顺序改先动models/下的网络定义再改options.py加参数最后在FedOur.py里挂上新的算法分支。# 以接入一个自定义 CNN 为例在 models/Nets.py 里加类 import torch.nn as nn class MyNet(nn.Module): def __init__(self, num_classes10): super().__init__() self.features nn.Sequential( nn.Conv2d(3, 32, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), ) self.classifier nn.Linear(64 * 8 * 8, num_classes) def forward(self, x): x self.features(x) x x.view(x.size(0), -1) return self.classifier(x)然后在options.py里加--backbone mynet在FedOur.py的模型构建处加一个分支判断。聚合逻辑不用大改只要保证所有客户端的state_dictkey 一致FedOur_Aggr.py里的平均操作就能直接用。验证聚合是否正确有个很土但有效的办法把--num_clients设成 2--fraction设成 1.0两个客户端用完全相同的数据和初始化跑一轮聚合后全局模型的参数应该等于两个客户端参数的算术平均。写个小脚本对比一下# 验证聚合结果是否等于简单平均 import torch client1 torch.load(client1.pth) client2 torch.load(client2.pth) global_model torch.load(global_after_aggr.pth) for key in client1: expected (client1[key] client2[key]) / 2 assert torch.allclose(global_model[key], expected, atol1e-6), fMismatch at {key} print(聚合验证通过)这个检查能帮你排除掉加权平均权重写错、聚合时漏掉某些层、或者客户端参数没同步更新这类问题。Meta-Transfer 部分验证起来更麻烦常见做法是先在单客户端上跑内循环确认支持集损失下降再扩展到多客户端。从那以后我每次改完聚合逻辑都会先用两个客户端跑一轮做参数对比确认没问题再上完整实验。联邦学习实验的坑大多不在算法本身而在数据管道和参数同步这些工程细节上把这两块盯住复现曲线基本不会偏太远。希望帮到你。本文还有配套的精品资源点击获取