☰
联邦学习实验从零跑通:三个实验+源代码+模型+图片演示
2026/10/7 3:33:08 网站建设 项目流程

简介:本资源是一套基于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配置及md说明文档,整体约631KB,目录涵盖模型定义、数据采样、聚合与本地更新等模块。已有227人学习下载。读者可获得完整可运行的实验代码、ResNet等模型实现、训练结果可视化图表与清晰的工程结构,便于复现实验、修改算法或直接用于论文与答辩展示。

1. 联邦学习实验从零跑通:三个实验到底在验证什么

联邦学习这个词这两年被提得很多,但真正动手跑过一轮完整实验的人并不多。我见过太多人卡在第一步:环境装完,数据不知道怎么切,模型不知道怎么分发,最后只能对着论文里的架构图发呆。这个标题里的「三个实验+源代码+模型+图片演示」,本质上是一套可以本地复现的联邦学习最小闭环——它要解决的不是理论推导,而是让你亲眼看到:数据不出本地的情况下,模型到底怎么协同训练、效果差多少、坑在哪里。

适合谁看?如果你已经会写基本的 Python 训练脚本,用过 PyTorch 或 TensorFlow 跑过单机模型,但没碰过联邦场景,这篇就是给你铺路的。三个实验通常对应三种典型设定:同构数据下的基础联邦平均、异构数据下的非独立同分布挑战、以及通信轮次与精度的权衡。源代码和模型文件的意义在于,你不用从零造轮子,但必须理解每一行在干什么,否则换个数据集就翻车。图片演示则是验证手段——损失曲线、准确率对比、混淆矩阵,这些图能告诉你实验有没有真的跑对。

2. 三个实验的设定拆解:从联邦平均到非独立同分布

2.1 实验一:同构数据下的联邦平均基线

第一个实验通常是最干净的设定:所有客户端的数据独立同分布,每个客户端拿到的样本类别分布一致。这个实验的目的是建立基线——如果连这种理想情况都跑不出合理精度,后面的实验不用看了。

联邦平均的核心逻辑是:服务器下发全局模型,客户端各自用本地数据训练若干轮,上传模型参数(或梯度),服务器按样本量加权平均。听起来简单,但代码里最容易出错的是参数聚合的顺序和权重计算。

# 联邦平均的核心聚合逻辑 import torch def federated_averaging(global_model, client_models, client_sizes): """ global_model: 全局模型 client_models: 各客户端训练后的模型列表 client_sizes: 各客户端样本数量列表 """ total_samples = sum(client_sizes) global_dict = global_model.state_dict() # 初始化聚合缓存 for key in global_dict.keys(): global_dict[key] = torch.zeros_like(global_dict[key]) # 按样本量加权累加 for client_model, size in zip(client_models, client_sizes): client_dict = client_model.state_dict() weight = size / total_samples for key in global_dict.keys(): global_dict[key] += client_dict[key] * weight global_model.load_state_dict(global_dict) return global_model

这段代码的关键在weight = size / total_samples。很多初学者直接做等权平均,结果某个客户端只有几十条样本却和几千条样本的客户端话语权一样,全局模型直接跑偏。参数说明:client_sizes必须和client_models一一对应,顺序不能乱;torch.zeros_like初始化时要注意数据类型,如果模型里有整型 buffer(比如 BatchNorm 的 num_batches_tracked),直接乘浮点权重会报类型错误,常见做法是跳过非浮点参数或单独处理。

实验一跑完后,你应该看到全局模型的准确率随着通信轮次上升,最终接近集中式训练的效果,但通常低 1 到 3 个百分点。如果差距超过 5 个点,先检查数据划分是否真的同分布,再看客户端本地训练轮数是不是太少。

2.2 实验二:非独立同分布数据的挑战与修正

第二个实验把数据打乱,让每个客户端只包含部分类别,模拟真实场景中用户行为差异。这时候联邦平均会明显掉点,因为各客户端的本地模型会偏向自己见过的类别,聚合时相互抵消。

常见修正手段有三种:一是客户端本地训练时加入正则项,限制本地模型偏离全局模型太远;二是服务器端做动量更新,不直接替换全局参数;三是调整客户端采样策略,每轮只选部分客户端参与。

# 带近端项的本地训练损失 def local_train_with_proximal(model, global_model, dataloader, epochs, mu=0.01): """ mu: 近端项系数,控制本地模型与全局模型的偏离程度 """ optimizer = torch.optim.SGD(model.parameters(), lr=0.01) criterion = torch.nn.CrossEntropyLoss() for epoch in range(epochs): for data, target in dataloader: optimizer.zero_grad() output = model(data) loss = criterion(output, target) # 近端项:惩罚本地参数与全局参数的差异 prox_term = 0.0 for local_param, global_param in zip(model.parameters(), global_model.parameters()): prox_term += ((local_param - global_param) ** 2).sum() loss += (mu / 2) * prox_term loss.backward() optimizer.step() return model

mu的取值很关键:太大则本地模型学不动,太小则退化成普通联邦平均。我一般从 0.01 开始试,观察全局准确率曲线,如果震荡厉害就加到 0.1,如果几乎不涨就降到 0.001。注意global_model在这一轮中不能被更新,它的参数是固定的参考点。

实验二的图片演示通常会展示不同mu值下的准确率对比,以及客户端本地模型的类别偏向热力图。如果你跑出来的结果和演示图差距很大,先确认数据划分的随机种子是否一致——非独立同分布的划分方式对结果影响极大。

2.3 实验三:通信轮次与模型精度的权衡

第三个实验关注效率:联邦学习的通信成本往往比计算成本更贵。实验三通常会对比不同通信轮次下的精度,以及是否使用梯度压缩、量化等技巧。

# 梯度量化示例:将浮点梯度压缩为低比特表示 def quantize_gradient(gradient, bits=8): """ gradient: 原始浮点梯度张量 bits: 量化比特数 """ min_val = gradient.min() max_val = gradient.max() scale = (2 ** bits - 1) / (max_val - min_val) # 量化 quantized = torch.round((gradient - min_val) * scale) # 反量化 dequantized = quantized / scale + min_val return dequantized

这个量化函数是最简单的线性量化,实际用的时候要注意:min_val和max_val如果是标量,需要先做全局归约;如果逐层量化,每层单独算。量化带来的精度损失在低比特时非常明显,8 比特通常还能接受,4 比特以下就要配合误差反馈。

实验三的图片演示一般会画一条「通信轮次-准确率」曲线,以及「压缩率-准确率损失」曲线。你需要关注的是拐点:多少轮之后精度不再明显上升,压缩到什么程度精度开始崩。这个拐点因数据集和模型而异,没有万能参数。

3. 本地环境搭建与源代码运行步骤

3.1 Python 环境与依赖安装的避坑清单

联邦学习实验对版本比较敏感,尤其是 PyTorch 和 numpy 的兼容性。我一般用 conda 建独立环境,避免和系统里的包打架。

# 创建并激活环境 conda create -n fl_experiment python=3.8 conda activate fl_experiment # 安装核心依赖 pip install torch==1.12.0 torchvision==0.13.0 pip install numpy==1.21.0 pip install matplotlib==3.5.0 pip install scikit-learn==1.0.2

版本号不是随便写的:torch 1.12 和 numpy 1.21 搭配比较稳,再新的 numpy 可能和旧版 torch 的 C 扩展冲突。如果你用 GPU,注意 CUDA 版本要和 torch 对应,torch.cuda.is_available()返回 False 的话先查驱动。

常见翻车点:pip 和 conda 混用导致包路径混乱。要么全用 pip,要么全用 conda,别交替装。另外,Windows 下路径分隔符和 Linux 不同,源代码里如果有硬编码的/,在 Windows 上可能读不到文件,改成os.path.join更稳妥。

3.2 数据划分与客户端配置

源代码里通常有一个config.py或args.py,控制客户端数量、数据划分方式、本地训练轮数等。以 CIFAR-10 为例,同构划分就是随机均匀分给每个客户端,非独立同分布划分则按类别分组。

# 非独立同分布数据划分示例 import numpy as np def dirichlet_split(labels, num_clients, alpha=0.5): """ labels: 所有样本的标签数组 num_clients: 客户端数量 alpha: Dirichlet 分布参数,越小越不均匀 """ num_classes = len(np.unique(labels)) client_indices = [[] for _ in range(num_clients)] for cls in range(num_classes): cls_indices = np.where(labels == cls)[0] np.random.shuffle(cls_indices) # 用 Dirichlet 分布生成每个客户端分到的比例 proportions = np.random.dirichlet([alpha] * num_clients) proportions = (np.cumsum(proportions) * len(cls_indices)).astype(int)[:-1] split_indices = np.split(cls_indices, proportions) for client_id, indices in enumerate(split_indices): client_indices[client_id].extend(indices) return client_indices

alpha越小,客户端之间的数据分布差异越大。alpha=0.5是常用的中等非独立同分布设定,alpha=0.1则非常极端。跑实验二的时候,建议固定随机种子,否则每次划分不一样,结果没法对比。

3.3 模型保存与图片演示生成

源代码里的模型保存通常用torch.save,但要注意保存的是state_dict还是整个模型。保存整个模型在加载时会依赖原始类定义,换环境容易报错,推荐只存参数。

# 保存和加载模型参数 torch.save(global_model.state_dict(), 'global_model_round_100.pth') # 加载时先实例化模型结构 model = ResNet18(num_classes=10) model.load_state_dict(torch.load('global_model_round_100.pth')) model.eval()

图片演示一般用 matplotlib 生成,损失曲线和准确率曲线画在一起时,注意双 Y 轴的刻度对齐。如果横坐标太密集(比如 1000 轮全画出来),用plt.xticks间隔采样,或者直接画平滑后的曲线。

4. 联邦学习实验的避坑与排查记录

4.1 全局模型不收敛,损失震荡剧烈

现象:每轮聚合后全局损失忽高忽低,准确率不升反降。

原因:客户端本地训练轮数过多,本地模型过拟合,聚合时参数差异太大。或者学习率设置过高,客户端更新步长太大。

解决:把本地训练轮数从 5 降到 1 或 2,学习率从 0.01 降到 0.001。如果还震荡,检查数据划分是否极端非独立同分布,适当增大alpha。

4.2 客户端数量多时内存溢出

现象:跑 100 个客户端时程序崩溃,报 OOM。

原因:源代码可能一次性把所有客户端模型加载到内存再聚合。100 个 ResNet 参数量叠加,显存或内存直接爆掉。

解决:改成逐客户端训练、逐客户端聚合,不要保留所有客户端模型副本。聚合时用累加代替列表存储。

4.3 准确率曲线和演示图对不上

现象:自己跑出来的准确率比图片演示低很多。

原因:随机种子不同、数据划分方式不同、或者演示图用的是最佳轮次而非最终轮次。

解决:固定所有随机种子(numpy、torch、random),确认数据划分参数一致。如果演示图标注了「最佳准确率」,那就要在代码里加模型选择逻辑,而不是取最后一轮。

4.4 GPU 利用率低,训练速度慢

现象:GPU 占用率只有 20% 到 30%,大部分时间在等数据。

原因:客户端串行训练,每个客户端数据量小,GPU 还没热起来就结束了。或者 DataLoader 的num_workers设为 0。

解决:把num_workers调到 4 或 8,客户端训练改成并行(如果显存够)。但要注意,并行客户端会改变聚合顺序,结果可能和串行略有差异。

4.5 模型保存后加载报错,提示缺少 key

现象:load_state_dict报Missing key(s) in state_dict。

原因:保存时用了DataParallel或DistributedDataParallel,参数名多了module.前缀。

解决:保存时用model.module.state_dict(),或者加载时用OrderedDict去掉前缀。最省事的办法是保存前先model = model.module如果用了并行包装。

5. 进阶技巧:用滑动窗口滤波平滑联邦训练曲线

联邦学习的准确率曲线往往比集中式训练更抖,因为每轮聚合的客户端组合不同。如果你要拿曲线做汇报或论文图,原始曲线不太好看。我一般用滑动窗口滤波做后处理,但注意:滤波只用于可视化,不能用来篡改实验数据。

# 滑动窗口滤波平滑曲线 def moving_average(data, window_size=5): """ data: 原始准确率列表 window_size: 窗口大小,奇数效果更好 """ smoothed = [] half = window_size // 2 for i in range(len(data)): start = max(0, i - half) end = min(len(data), i + half + 1) smoothed.append(sum(data[start:end]) / (end - start)) return smoothed

window_size的选择有讲究:太小起不到平滑作用,太大则拐点被抹掉。我通常取总轮次的 1% 到 2%,比如 100 轮取 3 或 5。滤波后的曲线用于展示趋势,原始数据仍然要保留在日志里。

另一个技巧是验证联邦模型是否真的学到了东西:拿全局模型在单个客户端的本地测试集上评估,如果精度远低于全局测试集,说明模型对某些客户端过拟合了。这个检查能帮你发现非独立同分布场景下的公平性问题。

我自己的习惯是每跑完一个实验,先把原始日志和模型文件归档,再用脚本统一生成图片。这样换参数重跑时,对比的是同一套可视化流程,不会因为画图代码改动导致误判。联邦学习实验的坑大多不在算法本身,而在数据划分和聚合细节上,多跑几轮、多存几个检查点,比事后调参省事得多。希望帮到你。

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询