ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

联邦学习实战:三个实验+源码+模型+图片演示,从FedAvg到非独立同分布

联邦学习实战:三个实验+源码+模型+图片演示,从FedAvg到非独立同分布 简介这是一份面向计算机、人工智能及相关专业学生与开发者的联邦学习实验资源围绕 FedAvg、FedPer、FedRep 与 FedOur 等算法展开对比研究适合课程设计、毕业设计、算法复现与进阶学习使用。资源包共43个文件包含14个Python源码文件、18张png与2张jpg实验曲线图、5个xml配置及说明文档压缩包约631KB代码结构涵盖本地更新、聚合、采样、模型定义与迁移模块并附ResNet18、ResNet34等网络实现。实验部分覆盖Cifar-10上多算法准确率与损失对比、MedMNIST在10/50/100客户端规模下的表现以及Chest X-Ray Images上FedAvg全局模型与本地模型、Meta-Transfer方案的比较图片直观展示训练趋势。目前已有227人学习下载后可参考完整源码、模型与实验图表快速理解联邦学习流程并在此基础上修改扩展。1. 从三个实验切入联邦学习到底在解决什么工程问题很多人第一次接触联邦学习是被数据不出本地还能联合建模这句话吸引的但真正动手时才发现难点根本不在概念而在怎么把三个实验跑通、怎么让模型在多方数据上不掉点。这个标题里的三个实验源代码模型图片演示本质上是一套可复现的最小验证闭环用 Python 把联邦学习的核心机制拆成三个递进实验每个实验都有独立源码、训练出的模型文件和可视化结果。它适合两类人——一类是想入门联邦学习但被论文公式劝退的工程师另一类是要在项目里快速验证联邦方案到底比集中式差多少的算法同学。我一般建议先别急着看杨强那本《联邦学习》的理论推导先把三个实验跑出图再回头补数学效率高得多。下面这套路径就是围绕能跑、能改、能对比来展开的。2. 三个实验的设计逻辑从横向联邦到非独立同分布2.1 为什么是三个实验而不是一个单个实验只能证明代码能跑证明不了联邦学习有效。三个实验的设计通常遵循一条递进线第一个实验验证基础流程即多个客户端各自持有数据、本地训练、上传参数、服务器聚合跑通 FedAvg 的最小闭环第二个实验引入数据异构让各客户端的数据分布不同观察模型精度如何下降这是联邦学习最真实的痛点第三个实验加入对比基线比如集中式训练和本地独立训练用图片把三条曲线画在一起直观回答联邦到底值不值。这种设计的好处是每个实验只增加一个变量排错时能快速定位是哪一层出了问题。常见做法是第一个实验用 MNIST 或 CIFAR-10 这类干净数据集第二个实验用 Dirichlet 分布切分数据模拟非独立同分布第三个实验固定随机种子做公平对比。2.2 实验一FedAvg 最小闭环的代码骨架下面这段代码是联邦学习最核心的聚合逻辑我把它压缩到能一眼看懂的程度。实际项目里会拆成 client.py、server.py、model.py 三个文件但核心就是这几步。import torch import torch.nn as nn import copy # 简单的两层全连接模型作为联邦学习的全局模型 class SimpleModel(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(784, 128) self.fc2 nn.Linear(128, 10) def forward(self, x): x torch.relu(self.fc1(x)) return self.fc2(x) def federated_averaging(global_model, client_models, client_sizes): FedAvg 核心按各客户端样本数加权平均参数 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], dtypetorch.float32) # 加权累加每个客户端的参数 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].float() * weight global_model.load_state_dict(global_dict) return global_model # 模拟 5 个客户端每个客户端本地训练 1 个 epoch global_model SimpleModel() client_models [] client_sizes [200, 300, 150, 250, 100] # 各客户端数据量不同 for i in range(5): local_model copy.deepcopy(global_model) # 这里省略本地训练循环实际会调用 local_train(local_model, dataloader) client_models.append(local_model) global_model federated_averaging(global_model, client_models, client_sizes) print(聚合完成全局模型已更新)这段代码的关键在federated_averaging函数它没有用简单的算术平均而是按size / total_samples加权这是 FedAvg 论文里的标准做法。参数client_sizes必须和client_models一一对应顺序错了聚合结果就完全不对。另一个容易翻车的点是torch.zeros_like后面要显式指定dtypetorch.float32否则在某些 PyTorch 版本上会默认成整型累加时直接截断小数。2.3 实验二用 Dirichlet 分布制造非独立同分布数据非独立同分布是联邦学习绕不开的坎。实验二的核心操作是把训练集按标签用 Dirichlet 分布切给各客户端让每个客户端只拿到部分类别的数据。下面这个函数我用了很多次参数alpha越小数据分布越偏。import numpy as np def dirichlet_split(labels, num_clients, alpha0.5): 按 Dirichlet 分布将数据索引分配给各客户端 labels: 全部训练数据的标签数组 num_clients: 客户端数量 alpha: 浓度参数越小分布越不均匀 返回: 每个客户端的数据索引列表 num_classes len(np.unique(labels)) client_indices [[] for _ in range(num_clients)] for c in range(num_classes): # 取出当前类别的所有索引 class_indices np.where(labels c)[0] np.random.shuffle(class_indices) # 用 Dirichlet 分布生成分配比例 proportions np.random.dirichlet([alpha] * num_clients) # 按比例切分最后一段用剩余量补齐避免舍入丢数据 split_points (np.cumsum(proportions) * len(class_indices)).astype(int)[:-1] splits np.split(class_indices, split_points) for client_id, split in enumerate(splits): client_indices[client_id].extend(split.tolist()) return client_indices # 使用示例假设有 60000 条训练数据10 个类别10 个客户端 # labels 是长度为 60000 的标签数组 # client_data dirichlet_split(labels, num_clients10, alpha0.5)alpha0.5是一个常用起点alpha0.1时每个客户端可能只有一两个类别的数据模型会严重偏向本地类别。这里有个血泪经验np.split的切分点必须用cumsum后取整再去掉最后一个否则最后一段会多出或漏掉数据。跑完这个实验后你会看到全局模型在客户端本地测试集上的精度波动明显变大图片演示里通常用箱线图来展示这种波动。2.4 实验三集中式、联邦、本地独立训练的对比实验三的价值在于给出一个可量化的结论。我一般会固定三组实验的超参数相同的模型结构、相同的学习率、相同的 batch size唯一变量是训练方式。集中式训练把所有数据拼在一起跑联邦训练走 FedAvg 流程本地独立训练则是每个客户端只用自己的数据训练且不聚合。对比结果通常呈现这样的规律集中式精度最高联邦次之本地独立最差且方差最大。如果联邦和集中式的差距在 2 个百分点以内这个联邦方案就值得继续投入如果差距超过 10 个百分点就要检查数据异构程度和聚合轮数。图片演示里一般会画三条折线横轴是通信轮数纵轴是测试精度联邦那条线通常在前几轮上升很快后面逐渐逼近集中式但始终有差距。3. 环境搭建与源码运行从零把三个实验跑起来3.1 Python 环境与依赖安装的确定性步骤联邦学习实验对版本比较敏感我建议用 conda 建独立环境别在系统 Python 里直接装。下面这套命令在 Linux 和 macOS 上都能跑Windows 用户把source activate换成conda activate即可。# 创建独立环境指定 Python 3.9这个版本对 PyTorch 兼容性最好 conda create -n fl_experiment python3.9 -y conda activate fl_experiment # 安装 PyTorchCPU 版本足够跑通三个实验有 GPU 的换对应 CUDA 版本 pip install torch1.13.1 torchvision0.14.1 --index-url https://download.pytorch.org/whl/cpu # 安装实验依赖 pip install numpy matplotlib scipy tqdm # 验证安装 python -c import torch; print(torch.__version__)这里指定torch1.13.1是因为这个版本在state_dict的 dtype 处理上比较稳定不会出现前面提到的整型截断问题。matplotlib用来画对比曲线tqdm用来显示训练进度条scipy在部分数据切分脚本里会用到。如果你用 vscode 配置 Python 环境记得把解释器选到fl_experiment这个 conda 环境否则会跑到系统 Python 上报一堆找不到模块的错。3.2 源码目录结构与运行顺序拿到源代码后别急着python main.py先看清楚目录结构。典型的三个实验源码会这样组织federated_learning_experiments/ ├── config.py # 超参数配置学习率、轮数、客户端数都在这里 ├── models.py # 模型定义三个实验共用 ├── client.py # 客户端本地训练逻辑 ├── server.py # 服务器聚合逻辑 ├── data_utils.py # 数据加载和 Dirichlet 切分 ├── exp1_fedavg.py # 实验一入口 ├── exp2_non_iid.py # 实验二入口 ├── exp3_comparison.py # 实验三入口 └── results/ # 存放模型文件和图片运行顺序必须是先跑实验一确认基础流程通了再跑实验二和实验三。每个实验脚本跑完会在results/下生成.pt模型文件和.png对比图。如果实验一就报错先检查config.py里的num_clients和batch_size是否和你的机器内存匹配客户端数设成 100 而内存只有 8G 的话光加载数据就会爆。3.3 关键参数怎么调通信轮数、本地 epoch、学习率这三个参数直接决定实验能不能复现出合理结果。通信轮数global_rounds一般设 50 到 100太少模型没收敛太多浪费时间本地 epochlocal_epochs设 1 到 5设太大客户端会过拟合本地数据聚合后反而掉点学习率lr联邦场景下通常比集中式小我一般从 0.01 起步观察前 10 轮 loss 曲线如果震荡就降到 0.005。参数推荐范围作用调参信号global_rounds50-100通信轮数loss 不再下降即可停local_epochs1-5本地训练轮数超过 5 容易过拟合lr0.005-0.01学习率loss 震荡则调小batch_size32-64批大小显存不足则调小alpha0.1-0.5Dirichlet 浓度越小异构越强这张表里的alpha只影响实验二但它决定了实验二能不能复现出精度下降的现象。如果alpha设成 10数据接近均匀分布实验二和实验一结果几乎一样就失去了验证意义。4. 避坑与排查联邦学习实验里最容易翻车的五件事4.1 聚合后精度不升反降现象每轮聚合后全局模型精度比上一轮还低曲线呈锯齿状。原因通常是客户端本地学习率过大导致本地模型跑偏太远聚合时把全局模型带偏。解决办法是把lr降到 0.005 以下或者减少local_epochs到 1让本地模型不要偏离全局太远。另一个可能是聚合时权重算错了检查client_sizes是否和实际数据量一致。4.2 非独立同分布实验里某些客户端精度为 0现象实验二中部分客户端在本地测试集上精度接近 0。原因是 Dirichlet 切分后这些客户端只拿到了极少数类别的数据模型在本地根本没学到其他类别。这其实是正常现象不是 bug。解决办法是在图片演示里用箱线图展示分布而不是只看平均值如果想让每个客户端至少有两个类别把alpha调到 0.5 以上。4.3 模型文件加载时报 key 不匹配现象load_state_dict报Missing key(s)或Unexpected key(s)。原因是保存模型时用了torch.save(model)整个对象加载时模型类定义变了或者用了DataParallel包装导致 key 多了module.前缀。解决办法是统一用torch.save(model.state_dict(), path)保存加载时先实例化模型再load_state_dict不要直接torch.load整个模型。4.4 图片演示中文乱码现象matplotlib 画出的图里中文变成方框。原因是默认字体不支持中文。解决办法是在画图脚本开头加两行配置import matplotlib.pyplot as plt plt.rcParams[font.sans-serif] [SimHei] # 指定中文字体 plt.rcParams[axes.unicode_minus] False # 正常显示负号Linux 上没有 SimHei 的话换成WenQuanYi Micro Hei或直接下载字体文件放到 matplotlib 的字体目录。4.5 实验三对比不公平现象联邦精度远低于集中式怀疑联邦方案不行。原因往往是集中式训练用了更多 epoch 或更大 batch而联邦受限于通信轮数。解决办法是固定总训练量集中式的 epoch 数乘以数据量应该约等于联邦的global_rounds × local_epochs × 单客户端数据量。只有总计算量对齐了对比才有意义。5. 进阶技巧用灾难性遗忘视角重新审视联邦聚合跑通三个实验后如果你还想再深挖一层我建议关注联邦学习里的灾难性遗忘问题。这个概念在热词里也出现过它指的是模型在学习新任务时忘记旧任务。在联邦场景下如果各客户端的数据分布随时间变化全局模型在聚合后可能对早期参与客户端的表现变差。一个实用的验证方法是在实验二的基础上让客户端分批次参与训练第一批训练完后记录精度第二批加入后再测第一批客户端的精度如果明显下降就说明存在遗忘。应对思路有两种。一种是服务器端保留历史全局模型的参数聚合时按一定比例混入旧参数相当于给模型吃后悔药另一种是客户端本地训练时加入正则项约束参数不要偏离全局模型太远。下面这段代码演示了服务器端参数混入的做法def aggregate_with_memory(global_model, client_models, client_sizes, prev_global, beta0.1): 带历史记忆的聚合新全局参数 (1-beta) * 联邦聚合结果 beta * 上一轮全局参数 beta 越大对历史记忆保留越多适合客户端数据分布变化快的场景 total sum(client_sizes) new_dict global_model.state_dict() prev_dict prev_global.state_dict() for key in new_dict.keys(): new_dict[key] torch.zeros_like(new_dict[key], dtypetorch.float32) for model, size in zip(client_models, client_sizes): w size / total for key in new_dict.keys(): new_dict[key] model.state_dict()[key].float() * w # 混入历史参数 for key in new_dict.keys(): new_dict[key] (1 - beta) * new_dict[key] beta * prev_dict[key].float() global_model.load_state_dict(new_dict) return global_modelbeta是关键参数设 0.1 表示保留 10% 的历史记忆设太大模型会停滞不前。这个技巧在客户端数据分布稳定的场景下收益不明显但在数据分布漂移的场景下能明显缓解精度回退。我自己踩过的坑是beta一开始设了 0.5结果模型十轮都不动后来降到 0.1 才找到平衡点。验证方法很简单跑实验二的同时记录每轮全局模型在最早参与客户端上的精度画成曲线看有没有明显下滑。如果没有下滑说明你的场景暂时不需要这个技巧如果有再调beta。这套三个实验的框架我前后改过七八版最大的体会是别一上来就追求复杂模型和花哨的聚合算法先把 FedAvg 在非独立同分布下的表现摸清楚再决定要不要上改进方案。很多项目翻车不是因为算法不够新而是因为数据切分和参数对齐这些基础工作没做扎实。希望帮到你。本文还有配套的精品资源点击获取
返回列表