尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyTorch实现FedAvg联邦学习:MNIST数据Non-IID切分与聚合实战

PyTorch实现FedAvg联邦学习:MNIST数据Non-IID切分与聚合实战 简介一份基于PyTorch的MNIST联邦学习实战代码面向机器学习初学者、算法研究者以及需要了解隐私保护训练方式的开发者演示FedAvg算法的完整工作流程。项目围绕手写数字识别任务展开通过本地客户端多次梯度更新和服务端参数聚合来迭代全局模型能够直观呈现分布式训练的协作过程。压缩包共17个文件包含8个Python脚本分别承担数据加载、模型定义、客户端本地训练、服务端协调与参数聚合等功能另有4个gz格式的MNIST数据文件、说明文档、备份文件及附赠内容总体约20.54MB目录结构清晰方便按模块对照学习已有134人学习浏览。通过这份代码读者可以深入理解联邦学习处理非独立同分布数据时的通信效率、模型收敛特性以及隐私保护优势同时PyTorch实现也为后续替换模型结构、接入不同数据集或扩展多客户端模拟提供了可操作的基础。代码注释清晰模块划分明确便于二次开发与实验参数调整。1. MNIST联邦学习代码FedAvg master为什么它是入门联邦学习最划算的一课这套项目名拆开看就是三件事用MNIST当实验场、用FedAvg当聚合算法、把服务端当成master来统一调度全局模型。它的价值不在把MNIST刷到多高精度——那本来就是90年代末就解决的问题——而在于让你在一台没有GPU的笔记本上也能把「数据不集中、模型轮流训练、参数聚合下发」这条链路完整跑通让联邦学习从论文公式变成你能肉眼观察的训练日志。适合谁看刚读完杨强老师的联邦学习综述、准备做课程作业或毕设开题的学生需要在Non-IID切分、客户端参与率、本地epoch这些参数上做对比实验的研究生以及想验证FedAvg基线、后续打算替换成FedProx或加差分隐私的算法工程师。它解决的核心痛点是论文里的FedAvg步骤看懂了一复现就各种意外——下载404、聚合后精度暴跌、结果对不上。这份代码把最常见的翻车点提前踩平了。2. FedAvg聚合公式拆解客户端平均为什么能逼近集中式训练2.1 联邦学习的三个角色服务端master、客户端worker、数据不出端联邦学习与传统分布式训练的最大区别是数据不动。传统分布式训练先把数据集按batch切好分给多张卡每个worker拿到的数据本身就是整体的一部分本质上还是把数据搬到了算力旁边。联邦学习反过来数据留在各自的客户端本地可能出现不同的分布服务端master只负责下发模型参数、收集更新、做聚合永远接触不到原始样本。在这个架构里master不是git分支名而是指服务端节点。它在每个通信轮次做三件事把当前全局模型广播给选中的客户端等待客户端本地训练完成后回收权重或梯度按规则聚合成新的全局模型。客户端worker则是持有本地数据的一方它的训练逻辑和普通模型训练没有本质区别区别在于它的训练起点是master下发的全局模型训练完回传的是模型更新而不是原始数据。这种「数据不出端、参数上云、模型下发」的设定让联邦学习在医疗、金融这类数据敏感的行业里格外受关注。MNIST虽然本身没有隐私诉求但它数据结构简单、样本量大、类别均衡用来验证联邦链路是否正确成本极低。跑通MNIST之后再换真实业务数据你只需要换数据集和模型结构聚合和服务端调度代码几乎不用动。2.2 FedAvg的聚合公式按样本量加权平均而不是简单求平均FedAvg的核心是一个加权平均公式。假设有K个客户端参与本轮聚合第k个客户端持有n_k条样本全局总样本数n等于所有n_k之和客户端本地训练结束后更新出的模型权重记为w_k那么全局模型w按下式更新w Σ (n_k / n) × w_k也就是说数据量大的客户端对全局模型的发言权更大。这个加权系数是FedAvg和朴素联邦平均最本质的区别。如果忽略样本量直接对所有权重取算术平均数据量大的客户端贡献会被低估模型会偏向小样本客户端的学习方向。和更早的FedSGD做对比会更容易理解FedAvg的动机。FedSGD要求客户端把本地梯度回传给服务端服务端对梯度做平均后再更新全局模型通信量和训练步数强绑定。FedAvg则让客户端本地跑多个epoch然后回传的是更新后的模型权重服务端做的是权重平均。它把一个通信轮次里的本地计算量放大通信频次降下来但精度损失很小。为什么权重平均能近似等于集中式SGD的效果直观解释是当客户端数据分布一致IID时每个客户端本地SGD的方向和全局梯度方向大致一致多走几步只是把这个方向走得更远一些所有客户端走完再平均等价于在全局梯度方向上迈了一大步。这和分布式训练里的梯度AllReduce在数学上是相近的。当客户端数据变成Non-IID时这个等价性被打破后面会看到这是精度掉落的根源。2.3 为什么MNIST是FedAvg最合适的Hello WorldMNIST在联邦学习里的地位相当于快速排序在算法课里的地位——样本足够多、类别足够清晰、训练足够快。6万张训练图28×28灰度最朴素的CNN在CPU上十几秒就能跑完一个epoch这让你可以在一个下午试完几十组参数组合而这种试错密度在CIFAR-10或者医疗影像上是做不到的。MNIST还有一个常被忽略的优点可视化极其方便。你可以把每个客户端拿到的数据切片打印成图片网格直观看到Non-IID切分后某个客户端是不是只剩两三个类别的样本。这种「看得见」的能力在调试联邦学习时非常重要因为聚合逻辑出错时精度数字只是表象数据分布才是根因。但也要提前打一针预防针MNIST太简单了简单到集中式训练随便就能做到99%以上的测试精度。同样的模型放到联邦学习里IID切分下做到98%以上不稀奇Non-IID强切分下掉到90%上下也不代表算法有问题。你的关注点应该放在「不同参数设置下精度如何变化」而不是和MNIST榜单上的SOTA对比。联邦学习实验的价值在相对差异不在绝对数字。3. 用PyTorch复现FedAvg最小代码从数据切分到聚合更新3.1 数据准备torchvision下载MNIST报404的本地回退方案第一步是拿到MNIST数据。常见做法是直接用torchvision的datasets.MNIST接口但这里有个真实存在且高频出现的问题——MNIST原始文件托管在第三方对象存储上torchvision内置的下载地址有时会返回404表现为下载到一半连接断开或者直接抛HTTPError。import torch from torchvision import datasets, transforms def load_mnist(root./data, downloadTrue): transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)), ]) try: train_set datasets.MNIST(root, trainTrue, downloaddownload, transformtransform) test_set datasets.MNIST(root, trainFalse, downloaddownload, transformtransform) except Exception: # 在线下载失败时回退到本地已有文件 train_set datasets.MNIST(root, trainTrue, downloadFalse, transformtransform) test_set datasets.MNIST(root, trainFalse, downloadFalse, transformtransform) return train_set, test_set这段代码的逻辑是先尝试在线下载一旦抛异常就改用downloadFalse从本地磁盘载入。前提是你已经手动把MNIST的四个.gz文件放到了./data/MNIST/raw/目录下目录结构和torchvision期望的一致。MNIST文件很小四个压缩包加起来约10MB手动下载后放到本地是成本最低的避坑手段。参数说明Normalize((0.1307,), (0.3081,))是MNIST全集的均值和标准差PyTorch官方范例沿用多年的固定统计量不要自己重新计算替换。downloadTrue只在首次运行使用后续建议固定为False避免每次启动都去检查网络。异常处理不要吞掉具体错误信息排查时把exception内容打印出来能省很多时间。3.2 Non-IID数据切分按标签用Dirichlet分布分给10个客户端联邦学习实验里数据切分方式直接决定实验结论。完全随机切分IID只能验证链路正确性论文里真正关心的是Non-IID——每个客户端手里的数据分布各不相同。常见做法是采用Dirichlet分布做标签级采样通过一个alpha参数控制异质强度这也是当前学术界的标准做法。import numpy as np from torch.utils.data import Subset def split_non_iid(train_set, client_num10, alpha0.5, seed42): labels np.array([train_set[i][1] for i in range(len(train_set))]) client_indices [[] for _ in range(client_num)] # 按数字0-9把样本索引分组 label_indices [np.where(labels d)[0] for d in range(10)] rng np.random.default_rng(seed) for label, idxs in enumerate(label_indices): rng.shuffle(idxs) # Dirichlet分布决定当前类别分到各客户端的比例 proportions rng.dirichlet(np.repeat(alpha, client_num)) splits (proportions * len(idxs)).astype(int) # 修正浮点取整导致的样本数偏差 splits[-1] len(idxs) - splits[:-1].sum() start 0 for cid, size in enumerate(splits): client_indices[cid].extend(idxs[start:start size].tolist()) start size return [Subset(train_set, idxs) for idxs in client_indices], label_indices逻辑拆解先按标签把全部样本索引分成10组然后对每个标签组独立做一次Dirichlet采样得到该类别分配到10个客户端的比例。alpha决定了分布的集中程度——alpha越接近0某个标签的样本越可能大量集中到少数几个客户端手里alpha取很大值比如100时每个客户端拿到的类别比例趋近均匀退化成近似IID。最后用Subset包装索引既保留原数据集又不复制像素数据内存友好。参数说明client_num10是论文里的标准客户端数量alpha0.5目前是复现Non-IID常用的中等强度设定我一般会再配一个alpha100做IID对照组。seed必须固定否则每次运行客户端分布都不同实验结果不可比。splits[-1]的修正行容易被人忽略浮点乘法和类型转换加起来可能让最后的客户端少几条或多几条数据必须补上这个差值。3.3 客户端本地训练定义模型和本地SGD更新客户端本地训练的网络结构不需要特殊设计MNIST上一个小型CNN就够用。注意FedAvg的客户端更新和普通训练的差异在「返回什么」——FedSGD返回梯度FedAvg返回更新后的完整state_dict这是两者在代码层面的关键分野。import torch.nn as nn import torch.nn.functional as F class SmallCNN(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(1, 32, 5, padding2) self.conv2 nn.Conv2d(32, 64, 5, padding2) self.fc1 nn.Linear(64 * 7 * 7, 512) self.fc2 nn.Linear(512, 10) def forward(self, x): x F.max_pool2d(F.relu(self.conv1(x)), 2) x F.max_pool2d(F.relu(self.conv2(x)), 2) x x.view(x.size(0), -1) x F.relu(self.fc1(x)) return self.fc2(x) def local_train(model, dataloader, epochs5, lr0.01, devicecpu): model.train() optimizer torch.optim.SGD(model.parameters(), lrlr, momentum0.9) criterion nn.CrossEntropyLoss() for _ in range(epochs): for x, y in dataloader: x, y x.to(device), y.to(device) optimizer.zero_grad() loss criterion(model(x), y) loss.backward() optimizer.step() return model.state_dict()逻辑说明网络是两个卷积层加两个全连接层参数量约82万CPU单机训练一个epoch大约十几秒适合需要反复跑参数组合的实验场景。local_train内部的训练循环看起来和普通训练完全一样但注意它没有返回loss只返回模型权重——因为服务端聚合时只需要权重不需要知道客户端训练过程的好坏。参数说明epochs5是FedAvg论文的默认本地轮数表示每个通信轮次里客户端要把本地数据完整训练5遍。epochs1时本方法退化成接近FedSGD的行为epochs过大时本地模型会偏移太远。lr0.01配合momentum0.9是MNIST上非常稳定的组合不建议在这个基础上把lr调大。客户端每次训练都要新建一个干净的模型实例并加载全局权重不要在客户端之间共享同一个模型对象否则你会踩到权重累积污染的坑。3.4 服务端master聚合更新按样本量加权平均与主循环聚合是FedAvg的灵魂。服务端收到客户端回传的state_dict后需要按每个客户端持有的样本量做加权平均然后把聚合结果写回全局模型。这里最容易写错的地方是直接在原模型上累加而不是先除以总样本数导致权重爆炸。def fedavg_aggregate(global_model, client_weights_list, client_sizes, total_size): global_dict global_model.state_dict() for key in global_dict: # 按样本量比例加权求和 global_dict[key] sum( (size / total_size) * w[key].float() for w, size in zip(client_weights_list, client_sizes) ) global_model.load_state_dict(global_dict) return global_model # 服务端主循环25个通信轮次每轮随机抽5个客户端 for rnd in range(25): need_clients np.random.choice(client_indices, size5, replaceFalse) client_weights, client_sizes [], [] for cid in need_clients: local_model SmallCNN() local_model.load_state_dict(global_model.state_dict()) w local_train(local_model, dataloader_dict[cid], epochs5) client_weights.append(w) client_sizes.append(len(datasets_list[cid])) global_model fedavg_aggregate( global_model, client_weights, client_sizes, sum(client_sizes) ) acc evaluate(global_model, test_loader) print(fround {rnd:02d}, acc{acc:.4f})逻辑说明聚合函数遍历全局state_dict的每个key对每个key做一次加权求和。size / total_size是分配到当前客户端的权重系数再乘上对应的权重张量。注意我用了.float()做类型转换防止CPU上float32和某些中间类型的隐式冲突。主循环里每轮随机抽5个客户端这对应参与率0.5——论文设置也是模拟「客户端不可能全部在线」的常态做法。参数说明replaceFalse表示每轮抽的客户端不重复但跨轮次允许重复。test_loader由3.1节里的test_set构建服务端持有测试集是允许的联邦学习的约束在训练数据测试集在服务端做全局验收是论文标准做法。如果你把测试集也切分下发到客户端评估口径就彻底乱了。聚合时还要保证所有客户端回传state_dict的key顺序完全一致同一份代码训练出的模型天然一致但如果你在某个分支改了网络结构就会踩坑。4. 把FedAvg跑起来命令入口、关键参数与收敛判据4.1 训练入口把所有实验变量暴露成命令行参数联邦学习实验的特性是参数多、组合多写死在代码里会让对比实验变成噩梦。常见的做法是用argparse把关键变量全部提出来每次跑实验只改命令行不碰代码。import argparse if __name__ __main__: parser argparse.ArgumentParser(descriptionMNIST FedAvg) parser.add_argument(--client_num, typeint, default10) parser.add_argument(--sample_ratio, typefloat, default0.5) parser.add_argument(--local_epoch, typeint, default5) parser.add_argument(--rounds, typeint, default25) parser.add_argument(--alpha, typefloat, default0.5) parser.add_argument(--seed, typeint, default42) args parser.parse_args()这六个参数基本覆盖了FedAvg复现实验的所有自由变量。client_num决定数据切分成多少份sample_ratio决定每轮实际参与训练的比例local_epoch是本地训练轮数rounds是全局通信轮次alpha控制Non-IID强度seed保证可复现。对应的运行命令是python main.py --client_num 10 --sample_ratio 0.5 --local_epoch 5 --rounds 25 --alpha 0.5 --seed 42把这行命令和下面的参数表对照你对这一轮实验在做什么应该一目了然。参数默认值作用调参方向client_num10客户端总数数据切分粒度sample_ratio0.5每轮参与比例越小噪声越大local_epoch5本地训练轮数越大通信成本越低但偏移越大rounds25全局通信轮次MNIST上10轮见趋势alpha0.5Non-IID强度越小数据分布越倾斜seed42随机种子对比实验必须固定这组默认参数在MNIST上能在CPU上约十几分钟跑完全部25个round。想快速验证代码链路有没有问题可以先跑5个round观察曲线趋势能跑通再把rounds拉到25。4.2 收敛判据观察精度曲线的形状而不是只看最后一个点很多初学复现的人只看最后一轮的测试精度这不够。FedAvg训练日志里最有信息量的是整条曲线的形态。前几轮精度快速上升说明链路正确、学习率合适5轮之后增速放缓是正常现象因为MNIST的简单分类边界已经被学得差不多了如果中段出现短暂的精度回落需要结合Non-IID强度判断——alpha0.1时曲线来回震荡是正常的alpha100时还震荡则说明学习率偏大。服务端评估函数如下每轮聚合后调用一次打印当前全局模型在完整测试集上的准确率。def evaluate(model, test_loader, devicecpu): model.eval() correct 0 with torch.no_grad(): for x, y in test_loader: x, y x.to(device), y.to(device) pred model(x).argmax(dim1) correct (pred y).sum().item() return correct / len(test_loader.dataset)这个函数没有新的数学重点在评估口径测试集必须完整保留在服务端不要分发到客户端。联邦学习的隐私约束是训练数据不上云测试数据在服务端做统一验收不违反这个设定。每轮打印的精度记录了全局模型随通信轮次的真实演化这才是判断收敛的依据。额外建议一个简单判据记录最后10个round的精度计算标准差。标准差低于0.5%说明实验已经稳定可以停掉高说明当前参数组合下模型还在剧烈波动需要检查lr或local_epoch。MNIST的简单性让这个判据非常可靠。4.3 和集中式训练对比算一算FedAvg到底「亏」了多少精度跑通FedAvg之后最有说服力的一张实验对比表是FedAvg和集中式训练在同等数据上的差距。集中式训练的做法是把全部6万张训练图合并用同一个SmallCNN、同样的SGD超参数训练同样的本地步数总和。下表给出MNIST上的典型结果区间不同seed下会有1%左右波动不要拿单次运行下结论实验设置客户端数据分布全局测试精度典型区间集中式训练所有数据合并99%以上FedAvg, alpha100近似IID98%左右FedAvg, alpha1.0轻度Non-IID95%左右FedAvg, alpha0.1强Non-IID88%到92%这个「精度梯度」本身说明两件事链路实现正确时FedAvg在IID下损失很小Non-IID程度越大精度掉得越多这是数据异质性带来的结构性代价不是代码bug。做实验报告时把这张表放进去比任何文字解释都有说服力。5. FedAvg复现避坑指南现象、定位与解决办法5.1 第一坑torchvision下载MNIST一直报404现象运行datasets.MNIST(downloadTrue)时抛HTTPError提示404 Not Found有时是下载到一半连接断开有时是torchvision缓存了损坏的临时文件导致重复失败。原因MNIST原始数据托管在AWS S3的某个公共桶里torchvision内置的URL并不总是可用。S3链接过期、网络策略变化、镜像不同步都可能触发404。这属于外部依赖失效不是你代码的问题。解决先手动下载train-images-idx3-ubyte.gz、train-labels-idx1-ubyte.gz、t10k-images-idx3-ubyte.gz、t10k-labels-idx1-ubyte.gz四个文件放进./data/MNIST/raw/目录保持文件名完全一致。然后在load_mnist里把download设为False直接读本地。3.1节给出的代码已经内含了异常回退逻辑手动放置文件后它会自动走本地路径。5.2 第二坑客户端本地精度很高全局模型精度却一直不涨现象每个客户端在本地数据上验证精度都能到92%以上但服务端在完整测试集上评估精度长期卡在80%左右甚至更低。你可能会怀疑聚合代码写错了。原因这是评估口径错位。客户端本地验证集来自Non-IID切分可能只含两三个数字类别模型在自己的数据切片上当然表现得很好。这个「本地精度」是严重偏估计完全不能代表全局效果。另一种常见诱因是客户端本地训练轮数太多模型已经过拟合到本地分布聚合后全局模型被平均得面目模糊。解决客户端训练时只记录loss不打印也不保存本地测试精度所有验收统一走服务端的test_loader。如果想把「本地精度」作为参考至少要确保它是随机切分的独立验证集而不是训练切片本身。Non-IID场景下只相信服务端评估结果。5.3 第三坑alpha调小后精度断崖式下跌爬不回来现象alpha0.1时前几个round精度还能到95%左右随后一路跌到88%以下后面无论跑多少轮都恢复不了。曲线呈现明显的「先涨后崩」形态。原因alpha0.1意味着每个客户端的本地数据高度集中极端情况下一个客户端只持有某一类数字。本地SGD会快速把模型推往自己类别的方向服务端把方向各异的权重平均后全局模型记住的只是「模糊的平均脸」。这是联邦学习里典型的灾难性遗忘——本地更新覆盖了之前轮次学到的全局知识聚合环节无法复原被冲掉的判别边界。解决先把alpha拉回0.5跑通确认链路没问题再往下压。如果业务场景必须强Non-IID优先降低local_epoch——从5降到2甚至1让本地更新别走太远。这两个参数组合起来调通常能救回几个点。如果仍然不满足就该换FedProx这类带近端项的算法但那已经不是这份FedAvg代码的范畴了。5.4 第四坑某个round之后loss变成NaN整个训练崩溃现象训练日志里前几轮都很正常突然某轮evaluate时报错或返回空值查看模型权重发现已经是NaN。原因两类常见诱因。一是学习率偏大某个客户端的本地训练直接发散发散后的权重被平均进全局模型污染全部后续轮次。二是客户端数据为空——比如alpha极小时某个客户端可能一条样本都没分到size / total_size计算出一个无意义值甚至除零。解决在聚合前对每个客户端权重做一次有限性检查发现异常直接丢弃该客户端。这个检查在生产型联邦系统里不是可选项而是必须项因为客户端是不可信的——可能是故障也可能是恶意。def validate_client_weight(w): for key, tensor in w.items(): if not torch.isfinite(tensor).all(): raise ValueError(fclient weight {key} contains NaN or Inf)5.5 第五坑同样代码换台机器跑结果对不上现象项目中记录的精度和另一个人复现出来的精度相差2个百分点以上两边都觉得自己跑的是对的。原因随机种子没固定是最常见的。数据切分、采样客户端、模型初始化、权重聚合里处处用到随机数不固定种子等于每次实验都是新实验。其次PyTorch在CPU上某些操作依赖底层数学库不同机器上的MKL或DNN实现会引入微小数值差异累积几轮后精度差1%很正常。解决在入口处同时固定torch.manual_seed、np.random.seed和random.seed并把DataLoader的num_workers固定为0或同一个值。MNIST实验通常做到「同机可复现」就够用了没必要跨机器逐位对齐。如果必须追求更强确定性可以开torch.use_deterministic_algorithms(True)但会牺牲性能小规模实验里性价比不高。6. 让MNIST上的FedAvg结果更可信两个验证技巧先说第一个技巧永远跑一组alpha100的IID对照组。这个组的客户端数据分布近似均匀理论上精度应该接近集中式训练的99%。如果这组也掉到95%以下说明问题出在实现链路上——聚合代码、数据切分、模型加载有bug反之如果这组正常、只有小alpha组精度低那才能证明你的Non-IID实验是有效的。我习惯把IID对照组的曲线注释保存在训练日志里每次改动代码后先跑这个组验证链路通过再跑正式实验组省掉大量排错时间。第二个技巧把每个round的日志落到结构化文件而不是只打印在控制台。一行CSV记录round编号、参与客户端ID列表、全局测试精度、平均本地loss攒够几十行后用脚本绘制精度曲线。曲线能直观暴露前面提到的先涨后崩、锯齿震荡等异常形态比盯着终端里的数字流可靠得多。实现上就是每次evaluate后追加一行with open(run_log.csv, a) as f: f.write(f{rnd},{args.alpha},{args.local_epoch},{acc:.4f}\n)这两个技巧的价值在跑正式实验时才会体现出来——当你跑了三组参数对比回头需要写报告或毕业论文时结构化日志和IID基线能直接变成成果的一部分而不是重新跑一轮。我自己做这份代码时吃过最大的亏就是一开始没留IID对照组某次改动让聚合权重系数写错连续三天在Non-IID里查精度为啥掉换成先跑基线、再上对照的思路之后问题半小时就定位了。希望这套流程对你也同样有效。本文还有配套的精品资源点击获取
返回列表