简介本资源是一套基于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, mu0.01): mu: 近端项系数控制本地模型与全局模型的偏离程度 optimizer torch.optim.SGD(model.parameters(), lr0.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 modelmu的取值很关键太大则本地模型学不动太小则退化成普通联邦平均。我一般从 0.01 开始试观察全局准确率曲线如果震荡厉害就加到 0.1如果几乎不涨就降到 0.001。注意global_model在这一轮中不能被更新它的参数是固定的参考点。实验二的图片演示通常会展示不同mu值下的准确率对比以及客户端本地模型的类别偏向热力图。如果你跑出来的结果和演示图差距很大先确认数据划分的随机种子是否一致——非独立同分布的划分方式对结果影响极大。2.3 实验三通信轮次与模型精度的权衡第三个实验关注效率联邦学习的通信成本往往比计算成本更贵。实验三通常会对比不同通信轮次下的精度以及是否使用梯度压缩、量化等技巧。# 梯度量化示例将浮点梯度压缩为低比特表示 def quantize_gradient(gradient, bits8): 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 python3.8 conda activate fl_experiment # 安装核心依赖 pip install torch1.12.0 torchvision0.13.0 pip install numpy1.21.0 pip install matplotlib3.5.0 pip install scikit-learn1.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, alpha0.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_indicesalpha越小客户端之间的数据分布差异越大。alpha0.5是常用的中等非独立同分布设定alpha0.1则非常极端。跑实验二的时候建议固定随机种子否则每次划分不一样结果没法对比。3.3 模型保存与图片演示生成源代码里的模型保存通常用torch.save但要注意保存的是state_dict还是整个模型。保存整个模型在加载时会依赖原始类定义换环境容易报错推荐只存参数。# 保存和加载模型参数 torch.save(global_model.state_dict(), global_model_round_100.pth) # 加载时先实例化模型结构 model ResNet18(num_classes10) 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_size5): 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 smoothedwindow_size的选择有讲究太小起不到平滑作用太大则拐点被抹掉。我通常取总轮次的 1% 到 2%比如 100 轮取 3 或 5。滤波后的曲线用于展示趋势原始数据仍然要保留在日志里。另一个技巧是验证联邦模型是否真的学到了东西拿全局模型在单个客户端的本地测试集上评估如果精度远低于全局测试集说明模型对某些客户端过拟合了。这个检查能帮你发现非独立同分布场景下的公平性问题。我自己的习惯是每跑完一个实验先把原始日志和模型文件归档再用脚本统一生成图片。这样换参数重跑时对比的是同一套可视化流程不会因为画图代码改动导致误判。联邦学习实验的坑大多不在算法本身而在数据划分和聚合细节上多跑几轮、多存几个检查点比事后调参省事得多。希望帮到你。本文还有配套的精品资源点击获取
阅读完成 · 觉得有帮助?