news 2026/9/24 18:59:38

联邦学习实战指南:三数据集算法对比与避坑实践

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
联邦学习实战指南:三数据集算法对比与避坑实践

简介:一份基于Python的联邦学习实验项目,面向人工智能、计算机及相关专业的学生、老师和开发者,既可作为课程设计、毕业设计的参考,也适合入门者理解联邦学习核心算法。项目包含三个递进实验:在Cifar-10上对比FedAvg、FedPer、FedRep与FedOur的准确率和目标损失;使用MedMNIST数据集测试不同客户端数量下的表现;基于Chest X-Ray Images数据集分析FedAvg全局模型、本地训练模型及Meta-Transf方法的效果差异,配套图片可直观查看训练曲线与结果。压缩包共43个文件,包括14个Python源码、18张PNG图表、2张JPG示意图、5个XML配置、Markdown说明与依赖清单等,文件类型覆盖代码、结果图表与项目配置,整体大小仅631KB,便于下载部署。所有代码均经测试运行成功,并附有实验图像,可复现三个实验并进一步修改参数;已有225人学习浏览,适合作为联邦学习方向的高分毕设或入门实践项目。

1. 联邦学习实验跑通三个数据集:这套源码到底能干什么

如果你是那种"论文读了十篇、代码一行没跑过"的联邦学习初学者,或者正在做毕设需要一组能对比 FedAvg、FedPer、FedRep 的完整实验,这套基于 Python 的联邦学习项目就是为你准备的。它把三组实验打包好了:Cifar-10 上四种算法对比、MedMNIST 上 10/50/100 客户端规模测试、Chest X-Ray 上的全局模型与元迁移对比,每个实验都生成 loss 和 acc 曲线图,想要输出论文插图或者毕设演示图,直接跑通就有素材。我不替你把话说满,但代码结构确实清晰,本地更新、聚合、采样、测试是分开的,改参数不用翻一整坨脚本。

2. 三个实验的比对逻辑:FedAvg 到 FedOur 的演进线

2.1 实验一:Cifar-10 上四种联邦学习算法的对比

Cifar-10 是计算机视觉里最常用的基准数据集之一,6 万张 32×32 彩色图片、10 个类别。在这个实验里,项目把 FedAvg、FedPer(Classify)、FedPer(Classify + 1 Block)、FedRep(Classify)和 FedOur 放在同一套数据划分和训练流程下对比,评价指标是准确率和目标损失值。

先说这几种算法的定位差异。FedAvg 是联邦学习的基线算法,McMahan 在 2017 年提出,思路是服务端下发全局模型,客户端本地训练几轮后只回传权重梯度或模型参数,服务端按样本量加权平均。FedPer 的全称是 Federated Personalization,它的核心观点是:神经网络可以拆成基底层(base layers)和个性化层(personalized layers),基底层共享,个性化层只在本地训练,不参与聚合。FedRep 类似,但侧重学习一个共享的数据表示(representation),再在每个客户端上训练分类头。至于 FedOur,这是项目自己的方法,你可以从 FedOur.py 和 FedOur_LocalUpdate.py、FedOur_Aggr.py 这三个文件里看到它的实现逻辑——我读下来的理解是,它在 FedPer 的拆分思路上做了扩展,对本地个性化层和全局共享层的更新频率、聚合权重做了更细的控制。

这个实验的价值在于:它能直接回答"在 Cifar-10 上,个性化联邦学习到底比 FedAvg 好多少"这个问题。跑完你会看到 FedAvg 的全局准确率通常比较稳定,但客户端本地泛化能力一般;FedPer 和 FedRep 在 Non-IID 数据分布下往往有更好的本地效果,但全局模型精度可能会略降。项目根目录里的 cifar-10-loss.png、cifar-10-acc.png、cifar-10-detail-loss.png、cifar-10-detail-acc.png 四张图,是已经跑好的结果,你可以拿它们和自己的输出对比。

2.2 实验二:MedMNIST 上客户端数量对收敛的影响

第二个实验换到了 MedMNIST——这是一个医学影像数据集家族,项目里实际用到的是其中两个子集:dermamnist(皮肤病变 7 分类)和 bloodmnist(血细胞 8 分类)。实验设计很直接:分别用 10、50、100 个客户端跑同一套联邦学习流程,观察客户端数量变化对最终精度和收敛速度的影响。

从联邦学习的理论看,客户端数量增加会带来两个效应。第一,每一轮参与聚合的客户端如果保持固定比例,那单轮通讯量增大,能更准确地估计全局梯度方向;第二,如果总训练轮数不变,单个客户端在每轮之间的本地数据被抽中的概率下降,数据覆盖变慢。项目生成的图片命名很直观:dermamnist_10clients_acc.png、dermamnist_50clients_loss.png、dermamnist_100clients_acc.png 等,loss 和 acc 是分开画的。我建议你看图时重点观察 100 客户端场景下训练初期的 loss 曲线是否比 10 客户端更陡峭——正常情况下应该更陡,因为每轮聚合见过的数据总量更大。

这个实验有个值得注意的细节:客户端数量改变时,默认的客户端采样比例和本地 epoch 数是否同步调整。如果你直接用默认参数把 10 客户端改成 100 客户端,很可能出现"客户端多了效果反而变差"的倒挂现象,这不是算法错了,而是每轮参与训练的客户端绝对数量翻了 10 倍,但总训练轮数没变,等效于每个客户端被访问的频次降低。跑这个实验前建议先看一眼 options.py 里的 num_clients、frac 和 local_epoch 三个参数的关系。

2.3 实验三:Chest X-Ray 上的全局模型与元迁移对比

第三个实验用的是 Chest X-Ray Images 数据集,肺炎 X 光胸片二分类。这一步的实验内容从项目摘要看是:对 FedAvg 的全局模型做 Local 训练得到的本地模型,和"我们的全局基本层经过 Meta-Transfer"之后的效果做对比。这里的关键词是 Meta-Transfer——元迁移学习。

元迁移的核心思路不是从零训练,而是让模型先学会"如何快速适应新任务"。具体到这个项目里,它应该是先把全局模型在源域数据上预训练好,然后冻结部分基本层,只对少量高层参数在目标数据上做快速微调,并用一种元学习的策略来更新。对应的代码文件是 transfer.py,配合 models/global_model.py 里的模型定义使用。你在运行时会看到 fine-tune-test-loss.jpg 和 fine-tune-test-acc.jpg 这两张输出图,展示的就是两种路线的收敛差异。

提示:第三个实验对显存的要求明显高于前两个。Chest X-Ray 原始图片分辨率远高于 Cifar-10 的 32×32,如果模型定义里没有全局池化层或者你改了输入尺寸,很容易把显存顶爆。

跑这个实验前建议先确认 dataset.py 里对 X 光图片的预处理管线:resize 到多少、有没有归一化、数据增强策略是什么。医学影像数据集的类别不均衡问题普遍存在,如果训练时直接用原始分布,模型很容易偏向多数类。

3. 把项目跑起来:从环境安装到输出第一张曲线图

3.1 环境准备:torch 版本与依赖的匹配

项目根目录有 requirements.txt,核心依赖就几样:PyTorch(含 torchvision)、NumPy、Matplotlib、Pillow。联邦学习代码本身不依赖任何特殊框架,用的是 PyTorch 原生的分布式通信——不是真的多机通信,而是单机模拟多客户端,客户端之间通过同一进程内的模型参数交换来模拟。

安装依赖时最容易踩的坑是 torch 和 torchvision 版本不匹配。我的习惯是先装 PyTorch,再根据它的版本来定 torchvision:

# 先创建虚拟环境,避免污染系统 Python python -m venv fedenv source fedenv/bin/activate # Windows 下用 fedenv\Scripts\activate # 安装 CPU 版或 CUDA 版 PyTorch,这里以 CUDA 11.8 为例 pip install torch==2.0.1 torchvision==0.15.1 --index-url https://download.pytorch.org/whl/cu118 # 再安装其他依赖 pip install -r requirements.txt

逻辑说明:PyTorch 和 torchvision 的版本绑定很紧,torch 2.0.1 对应 torchvision 0.15.1,混搭会在 import 时直接报错。requirements.txt 里如果没有精确锁版本,建议手动指定一组你已知能用的版本。--index-url参数指定了 CUDA 11.8 的 wheel 源,如果你机器上的 CUDA 版本不同,把 cu118 改成 cu117 或 cu121 即可;没有 N 卡就装 CPU 版,代码一样能跑,只是慢一些。

参数说明:虚拟环境这步不是可选项。这个项目依赖的 numpy 版本范围可能和你系统里其他项目冲突,虚拟环境能省掉大量"装了这个坏了那个"的问题。如果你用的是 Windows,注意 Python 版本建议 3.8 到 3.10,太新的 Python 3.12 在某些 torch 版本下没有预编译 wheel。

3.2 项目结构与三个实验的入口

解压 federal-learning-experiment-master.zip 后,你会看到如下核心文件:

文件作用
FedOur.py主入口,配置实验并启动训练流程
FedOur_LocalUpdate.py客户端本地更新逻辑,定义本地训练过程
FedOur_Aggr.py服务端聚合逻辑,定义参数聚合方式
options.py全部命令行参数的定义与默认值
dataset.py数据集的加载与预处理
sampling.py客户端数据划分,IID/Non-IID 采样策略
models/Nets.py小型网络定义
models/Resnet18.py、Resnet34.pyResNet 变体定义
models/global_model.py全局模型定义
transfer.py实验三的元迁移逻辑
test.py测试与评估脚本

运行第一个实验(Cifar-10 算法对比)的典型命令:

python FedOur.py --dataset cifar10 --model resnet18 --num_clients 20 --frac 0.5 --local_epoch 5 --rounds 50

逻辑说明:FedOur.py 接收命令行参数后,先按--dataset加载对应的数据集和预处理,再按--num_clients把数据切分给模拟客户端,--frac控制每轮参与训练的客户端比例。--model resnet18指定模型结构,服务端初始化全局模型后发给选中客户端,客户端本地跑--local_epoch个 epoch,回传参数,服务端用 FedOur_Aggr.py 里的聚合规则更新全局模型,循环--rounds轮。

参数说明:--frac 0.5配合--num_clients 20意味着每轮只有 10 个客户端参与训练。这不是随机抽样——sampling.py 里实现了参与客户端的轮换策略,保证各客户端被抽中的概率均衡。--local_epoch不宜设太大,联邦学习的本意是客户端只做少量本地更新,一般 1 到 10 之间;设太大会导致客户端模型漂移,聚合后全局模型反而变差。

3.3 输出物解读:loss 曲线、acc 曲线和模型检查点

训练结束后,根目录会生成一组 PNG 图片,命名规则是数据集_客户端数_acc/loss.png。Matplotlib 画图逻辑在 FedOur.py 或 utils 脚本里,每轮记录全局模型在测试集上的 loss 和 top-1 acc,最后统一出图。我拿到新环境跑通后,一般先看 loss 曲线的收敛形态:如果曲线在前 10 轮就快速下降然后趋于平缓,说明超参大致合理;如果 loss 全程不降或者震荡剧烈,优先检查学习率,联邦学习场景下全局学习率通常要比单机训练小一个数量级。

还有一个细节值得注意:训练过程中是否打印每轮耗时。联邦学习实验的一大痛点是慢——每轮要模拟多个客户端分别做前向反向传播,20 个客户端、每个 5 个 epoch,一轮训练可能就要几分钟。如果你的机器没有 GPU,建议先把--rounds调小到 10 先验证流程能走通,再跑完整实验。

4. 代码结构精读:六个核心脚本的职责与参数落点

4.1 客户端本地更新:FedOur_LocalUpdate.py 的黑匣子拆解

FedOur_LocalUpdate.py 是这套代码里最重要的单文件,它定义了客户端在收到全局模型后,如何在本地数据上做训练。这个文件通常是一个 LocalUpdate 类,构造函数接收全局模型参数、本地数据集、超参配置,核心方法 train 负责执行本地训练并返回更新后的模型参数。

class LocalUpdate(object): def __init__(self, args, dataset, idxs): # args: 全局参数配置对象 # dataset: 完整数据集 # idxs: 分配给该客户端的样本索引 self.args = args self.trainloader = DataLoader(DatasetSplit(dataset, idxs), batch_size=self.args.local_bs, shuffle=True) self.criterion = nn.CrossEntropyLoss() self.optimizer = torch.optim.SGD(self.model.parameters(), lr=self.args.lr, momentum=self.args.momentum) def train(self): for epoch in range(self.args.local_epoch): for images, labels in self.trainloader: images, labels = images.to(self.device), labels.to(self.device) self.optimizer.zero_grad() output = self.model(images) loss = self.criterion(output, labels) loss.backward() self.optimizer.step() return self.model.state_dict()

逻辑说明:这段代码是 FedAvg 系列算法客户端侧的通用骨架。DatasetSplit是一个 PyTorch Dataset 包装类,根据传入的idxs索引列表筛出属于该客户端的子集,从而在不复制原始数据的前提下实现数据划分。本地优化器用的是带动量的 SGD,state_dict()返回的是整个模型的参数字典,后续服务端聚合的就是这个字典。

参数说明:local_bs(本地 batch size)在联邦场景下有讲究。客户端数据量本来就少,batch size 太大可能导致每个 epoch 只有一两个 step,梯度更新次数不够;我一般设为 8 到 16 之间。lr是客户端本地学习率,服务端聚合时通常还会乘一个缩放系数,代码里可能在 FedOur_Aggr.py 中体现。

FedOur 和 FedAvg 在本地更新上的差别在于:FedOur 可能只让部分层参与本地训练,或者对本地训练后的参数做某种修正再返回。你可以在 train 方法里看到freeze_layers或梯度掩码相关的逻辑,这就是实验一的对比核心。

4.2 服务端聚合:FedOur_Aggr.py 与 FedAvg 的差异点

FedOur_Aggr.py 实现的是服务端把客户端回传的参数合并成新的全局模型。FedAvg 的标准做法是加权平均,权重是各客户端本地样本数占总样本数的比例:

def FedAvg(w, size): # w: 客户端参数权重列表, 格式 [{layer_name: tensor}, ...] # size: 各客户端的样本数量列表 total_size = sum(size) w_avg = copy.deepcopy(w[0]) for k in w_avg.keys(): w_avg[k] = w_avg[k] * size[0] / total_size for i in range(1, len(w)): w_avg[k] += w[i][k] * size[i] / total_size return w_avg

逻辑说明:这里的核心是按样本量加权。客户端本地数据量越大,它对全局模型的影响就应该越大,这是 FedAvg 的理论基础。但 FedOur 的聚合可能做得更细——比如对不同网络层使用不同的聚合权重,或者对个性化层不聚合、只聚合共享层。

跟 FedAvg 相比,FedOur 聚合时需要注意的坑是:如果你需要实现 FedPer 或 FedRep 的对比实验,不能简单地把全部参数都做聚合。FedPer 的个性化层在本地训练后应当保留本地版本,不参与服务端聚合;FedRep 则可能只聚合表示层。代码里 models 目录下分了 Nets.py、Resnet18.py、Resnet34.py、global_model.py 四个文件,就是在不同网络结构上做"哪些层共享、哪些层个性化"的切片实验。

4.3 options.py 与 sampling.py:参数基准和数据划分策略

options.py 是所有可调参数的中央枢纽。典型参数包括:--dataset(数据集选择)、--model(模型结构)、--num_clients(模拟客户端总数)、--frac(每轮参与比例)、--local_epoch(本地训练轮数)、--local_bs(本地 batch size)、--lr(学习率)、--rounds(全局通信轮数)、--iid(是否使用独立同分布数据划分)。

sampling.py 里的数据划分逻辑直接决定了实验的 Non-IID 程度。IID 划分是把数据随机打乱后均分给各客户端,每个客户端的类别分布基本一致,这种场景下联邦学习相对容易收敛。Non-IID 划分是按类别分块,比如 10 个类别分给 20 个客户端时,可以让每个客户端只拿 2 到 3 个类别的数据,模拟真实世界的"数据孤岛"。我跑实验时会故意在--iid参数下做两组对比,因为个性化联邦学习算法的优势恰恰在 Non-IID 场景下才能体现出来——如果数据是 IID 的,FedAvg 往往表现已经足够好,FedPer 的优势就不明显了。

提示:改采样逻辑前先把原版跑一遍,记录 baseline 精度。很多同学一上来就改 sampling.py 里的划分比例,结果模型不收敛,分不清是算法问题还是数据划分问题。

5. 避坑手册:跑这个项目最常见的四类翻车现场

5.1 现象:老代码爆 NumPy 兼容性错误,np.float不存在

你如果用的是 NumPy 1.24 以上版本,运行 dataset.py 时可能直接报AttributeError: module 'numpy' has no attribute 'float'。原因是 NumPy 1.20 之后移除了np.floatnp.int这些 Python 内置类型的别名,而不少老代码里还残留np.float的写法。解决方法是固定 NumPy 版本,或者改源码里的类型引用:

pip install numpy==1.23.5

如果不想降版本,把代码里的np.float替换为floatnp.int替换为int即可。这个坑几乎影响了所有 2021 年以前写的 PyTorch 项目,不是这个项目独有的问题。

5.2 现象:Cifar-10 或 MedMNIST 数据集下载卡死不动

第一次运行时,dataset.py 会自动下载数据集。Cifar-10 在国内网络环境下经常下载到一半断掉,torchvision 的下载逻辑没有断点续传,报错后本地留下一个损坏的压缩包,下次运行还会继续失败。解决方法是手动下载后放到指定位置:Cifar-10 放到data/cifar-10-batches-py/,MedMNIST 放到data/medmnist/下。项目 README 里一般会注明数据目录结构。另外要注意 MedMNIST 需要额外安装medmnist这个 Python 包,requirements.txt 里如果没有它就手动补上:

pip install medmnist

5.3 现象:显存不足,跑第一个实验就 OOM

ResNet18 在 32×32 的 Cifar-10 上参数量并不大,不该爆显存。但如果你的 batch size 太大,或者模型代码里没有适配低分辨率输入,比如 ResNet 默认的第一个卷积层 stride 和 pooling 会显著压缩特征图尺寸,在 32×32 输入下有些实现会直接报错或显存异常。解决方式是调小--local_bs,从 64 降到 16 试一下。另外确认一下代码里是否在模型定义后调用了.cuda()nn.DataParallel,多卡环境下 DataParallel 会默认占用所有可见 GPU,用CUDA_VISIBLE_DEVICES=0限定单卡:

CUDA_VISIBLE_DEVICES=0 python FedOur.py --dataset cifar10 --local_bs 16

5.4 现象:聚合后的全局模型精度反而低于单机训练的随机初始化模型

这不是 bug,而是联邦学习的典型现象。原因通常有三个:一是客户端本地学习率偏大,各客户端在本地"各自为政",参数更新方向互相抵消;二是参与客户端比例过低,比如 100 个客户端每轮只抽 1 个,全局模型基本上在随机游走;三是local_epoch设置过大,客户端过拟合到本地数据分布上。排查顺序是:先降低本地学习率到 0.01 以下,再提高--frac到 0.2 以上,最后把--local_epoch降到 1 试跑 20 轮看趋势。这个调试过程就是联邦学习里说的"调参玄学",但实际上它就这三个旋钮。

5.5 现象:实验二改客户端数量后,曲线对比没有规律

摘要里提到 MedMNIST 的客户端数量是 10、50、100 三组。如果你直接只改--num_clients而不动其他参数,10 客户端的实验可能比 100 客户端收敛更快、精度更高,但这不能说明"客户端越少越好"。因为 100 客户端时,每个客户端分到的样本量只有 10 客户端时的十分之一,本地训练数据严重不足。正确的做法是观察时固定每轮参与绝对数量,比如都是 10 个客户端参与,那--frac分别设为 1.0、0.2、0.1,让三组实验每轮见到的数据量一致,才能体现出客户端数量本身的效应。

6. 进阶验证:把精度复现从玄学变成可控的固定种子实验

复现联邦学习实验是一件比单机训练更麻烦的事,因为随机性来自三个层面:数据划分的随机性、客户端本地训练的随机性(Dropout、初始化等)、每轮参与客户端的抽样随机性。项目代码里如果没有在你主入口设置固定随机种子,每次跑出来的曲线都会有差异,你很难判断一个改动到底是真的有效还是随机波动。

我的习惯是写一个 set_seed 函数,在 main 的开头统一调用:

def set_seed(seed=42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False

逻辑说明:前三行管住了 Python 内置随机数、NumPy 随机数和 PyTorch CPU 随机数,manual_seed_all管住所有 GPU 上的随机数生成器。最后两行的关键作用是关闭 cuDNN 的自动调优算法选择——cuDNN 的 benchmark 模式会从多个卷积算法里挑最快的,不同轮次可能选到不同算法,导致结果不可复现,这在联邦学习的逐轮对比中是致命的。

参数说明:deterministic = True开启确定性算法会牺牲少量训练速度,但对于对比实验来说,这个代价完全值得。跑完整的三组客户端实验(10/50/100)时,我会把这三组不同客户端数的实验分别用不同的 seed 跑三次,取均值和方差画误差棒图,这样的图表在论文里才有说服力,不会被人质疑是挑了一次好运气的结果。

另外一个实际技巧是:每次实验启动时,把命令行参数连同时间戳一起存成一个 JSON 文件,放在输出目录里。这样一周后回来看某张曲线图,你还记得当初用的是哪组参数。如果 FedOur.py 本身没有记录参数的功能,就在运行命令前手动执行一下:

python FedOur.py --dataset medmnist --num_clients 50 --frac 0.2 --local_epoch 5 --rounds 30 --seed 42 2>&1 | tee run_$(date +%Y%m%d_%H%M%S).log

tee命令把训练过程同时输出到终端和日志文件,2>&1把 stderr 也合并进去,这样就算中途报错,也能完整回溯当时的环境。

这套项目最打动我的一点,是它把三个实验的产出物——loss 曲线、acc 曲线、微调对比图——全部留在了 img 目录里。我刚拿到代码时,先看了 dermamnist_100clients_loss.png 和 cifar-10-detail-acc.png 这两张图,心里对"曲线应该长什么样"有了底,再去跑自己的实验,如果输出和参考图形态差太远,就知道参数出了问题。从那以后,我每跑一个新的联邦学习项目,都会先找作者的输出图,再跑自己的复现——先确认终点长什么样,再决定要不要出发。希望这一套流程也能帮你在毕设或课设的实验环节里省下几个晚上的调试时间。如果你还是跑不通,或者想换自己的数据集,带着报错信息来,我把这十几个坑的排查顺序再帮你捋一遍。

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

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

仿真环境接入AI智能体:从API文档到Tool配置的实战指南

上个月我接到一个需求:把团队内部一直在跑的一个仿真环境(内部代号就叫 Sim)接入 AI 智能体,让模型可以直接通过自然语言调用 Sim 做场景验证。听起来不复杂,但真正动手才发现,从一份 API 文档到一段能用的…

作者头像 李华
网站建设 2026/9/24 18:58:33

Android横屏与多屏幕适配实战:从生命周期到刘海屏的完整指南

横屏适配做了这么多年,踩过的坑比踩过的键盘还多。每次产品经理一句“这个页面横屏看一下”,接下来的几小时基本就是和Activity重建、布局错乱、刘海屏遮挡斗智斗勇。而屏幕适配这件事,更是从早期的dimens多套文件到现在的sw限定符&#xff0…

作者头像 李华
网站建设 2026/9/24 18:58:29

红队渗透测试实战复盘:从入口突破到内网横向的完整攻击链拆解

红队测试这行干久了,你会发现一个有意思的现象:很多企业觉得自己的安全防护做得不错,等真正被红队模拟真实攻击者打一轮,往往撑不过两周。我印象最深的一次项目,目标是互联网上一家成熟的软件公司,防守方部…

作者头像 李华
网站建设 2026/9/24 18:57:50

Linux目录操作进阶:从opendir、readdir到scandir与nftw实战

处理过几百G日志目录,见过那种一个子系统一个目录、里面再按日期套时间戳的典型目录布局,就会明白Linux系统编程里的目录操作,从来不是表面看起来“打开目录读一遍”那么简单。前阵子我为了给一套跨NFS的日志系统写目录清理脚本,把…

作者头像 李华