简介一份基于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、FedPerClassify、FedPerClassify 1 Block、FedRepClassify和 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 torch2.0.1 torchvision0.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_sizeself.args.local_bs, shuffleTrue) self.criterion nn.CrossEntropyLoss() self.optimizer torch.optim.SGD(self.model.parameters(), lrself.args.lr, momentumself.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索引列表筛出属于该客户端的子集从而在不复制原始数据的前提下实现数据划分。本地优化器用的是带动量的 SGDstate_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.float、np.int这些 Python 内置类型的别名而不少老代码里还残留np.float的写法。解决方法是固定 NumPy 版本或者改源码里的类型引用pip install numpy1.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 medmnist5.3 现象显存不足跑第一个实验就 OOMResNet18 在 32×32 的 Cifar-10 上参数量并不大不该爆显存。但如果你的 batch size 太大或者模型代码里没有适配低分辨率输入比如 ResNet 默认的第一个卷积层 stride 和 pooling 会显著压缩特征图尺寸在 32×32 输入下有些实现会直接报错或显存异常。解决方式是调小--local_bs从 64 降到 16 试一下。另外确认一下代码里是否在模型定义后调用了.cuda()或nn.DataParallel多卡环境下 DataParallel 会默认占用所有可见 GPU用CUDA_VISIBLE_DEVICES0限定单卡CUDA_VISIBLE_DEVICES0 python FedOur.py --dataset cifar10 --local_bs 165.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(seed42): 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 21 | tee run_$(date %Y%m%d_%H%M%S).logtee命令把训练过程同时输出到终端和日志文件21把 stderr 也合并进去这样就算中途报错也能完整回溯当时的环境。这套项目最打动我的一点是它把三个实验的产出物——loss 曲线、acc 曲线、微调对比图——全部留在了 img 目录里。我刚拿到代码时先看了 dermamnist_100clients_loss.png 和 cifar-10-detail-acc.png 这两张图心里对曲线应该长什么样有了底再去跑自己的实验如果输出和参考图形态差太远就知道参数出了问题。从那以后我每跑一个新的联邦学习项目都会先找作者的输出图再跑自己的复现——先确认终点长什么样再决定要不要出发。希望这一套流程也能帮你在毕设或课设的实验环节里省下几个晚上的调试时间。如果你还是跑不通或者想换自己的数据集带着报错信息来我把这十几个坑的排查顺序再帮你捋一遍。本文还有配套的精品资源点击获取
