☰
PyTorch联邦学习实战:FedAvg算法实现MNIST手写数字识别
2026/9/28 16:37:50 网站建设 项目流程

简介:一套基于MNIST手写数字数据集的联邦平均算法(FedAvg)完整代码,采用PyTorch框架编写,面向机器学习初学者和算法研究人员,也适合希望在数据不出域前提下开展协作训练的开发者。资源包共包含17个文件,其中8个Python脚本用于数据加载、模型构建、服务端协调和客户端本地训练;4个压缩数据包存放MNIST图像数据;另有说明文档和附赠内容压缩包。整体大小约20.54MB,目录按功能划分,便于检索复用。目前已有134人学习。代码覆盖了联邦学习的关键流程:服务端负责接收并聚合客户端模型更新,客户端在本地完成多轮梯度下降,从而实现隐私保护下的分布式训练。FedAvg算法通过加权聚合与多轮本地迭代降低通信开销,在非独立同分布数据上依然有较好鲁棒性。附带脚本与备份资料还可作为课程设计、科研实验以及实际系统部署的参考起点。

1. 手写数字识别遇上联邦学习:这个 MNIST FedAvg 项目到底解决了什么问题

做过深度学习的人都知道 MNIST 手写数字识别,但绝大多数人是在单机环境下用完整数据集训练。联邦学习把游戏规则改了:数据不能出本地设备,模型却要协同训练。这套 FedAvg 代码把「数据不出域 + 模型全局共享」的矛盾用最经典的方式化解了——客户端各自用自己的本地数据训练,服务端只聚合模型参数,不碰原始样本。对于研究隐私保护、做医疗金融场景原型验证的从业者来说,这份代码是理解联邦学习落地细节的一个很好的起点。项目基于 PyTorch 实现,结构不复杂,但涉及数据划分、客户端训练、服务端聚合的完整闭环,适合作为二次开发的基线工程。


2. 先拆文件结构:拿到 FedAvg 工程后如何快速定位核心代码

2.1 压缩包的目录映射与模块职责

打开 FedAvg-master.zip,里面文件不多,但每个文件都有自己的角色。先理清职责再动手,避免在错误的地方浪费时间。

文件职责关键内容
server.py联邦服务端全局模型初始化、按轮调度客户端、FedAvg 聚合
clients.py联邦客户端本地训练、模型参数上传、接收全局参数
Models.py模型定义用于 MNIST 分类的神经网络结构
dataSets.py数据加载与切分MNIST 原始数据读取、iid/non-iid 数据划分
getData.py数据下载辅助下载/定位 MNIST 四个 gz 文件
README.md使用说明环境依赖、运行入口、参数说明
data/原始数据目录t10k-images-idx3-ubyte.gz 等四个文件
use_pytorch/框架标识目录确认本项目基于 PyTorch 实现
附赠内容.zip额外资源预处理后的数据划分或模型参数备份

从文件命名可以看出,这套代码刻意把「数据」「模型」「服务端」「客户端」拆成独立模块,这是联邦学习工程的常见组织方式。server.py 和 clients.py 是核心,dataSets.py 决定了数据怎么分——这一步直接影响实验结论的可信度。

2.2 从零跑通项目:按依赖顺序的执行路径

我习惯按「数据处理 → 单机验证 → 联邦联调」的顺序复现项目。先确认环境依赖:

pip install torch torchvision numpy

提示:PyTorch 版本建议 1.13 及以上,低版本在部分聚合操作上可能存在接口差异。

第一步先跑数据准备脚本:

python getData.py

这个脚本会检查 data/ 目录下是否存在 MNIST 的四个 gz 文件。如果文件已存在,直接加载;如果缺失,脚本尝试从网络下载。注意这里有一个常见的坑——torchvision 自带的下载接口经常遇到 404 问题,所以稳妥做法是手动把数据集放进去,这个在后面的避坑章节细说。

数据就绪后,单独验证客户端能否正常训练:

python clients.py

从代码里我们可以看到第 30 行附近的本地训练逻辑:

def local_train(model, train_loader, epochs, lr, device): model.train() optimizer = torch.optim.SGD(model.parameters(), lr=lr) for epoch in range(epochs): for batch_idx, (data, target) in enumerate(train_loader): data, target = data.to(device), target.to(device) optimizer.zero_grad() output = model(data) loss = torch.nn.functional.cross_entropy(output, target) loss.backward() optimizer.step() return model.state_dict()

这段代码的作用是让客户端用自己的本地数据完成多轮梯度下降,返回更新后的权重。epochs是本地训练轮数,lr是学习率,这两个参数直接决定模型收敛质量,后面的参数调优章节会展开。

单机验证通过后,启动联邦训练主入口:

python server.py

服务端会初始化全局模型,然后将模型参数分发给参与客户端,客户端在本地训练后将新参数返回,服务端执行 FedAvg 聚合更新全局模型,再进入下一轮。完整的联邦训练循环由 server.py 驱动,clients.py 只是被调用的组件。

2.3 数据加载的格式细节与 PyTorch 的对接方式

MNIST 原始数据不是常见的图片文件夹,而是四个 gz 压缩的二进制文件。dataSets.py 的核心工作是把这些二进制数据解析成 PyTorch 可以消费的张量。

import gzip import numpy as np import torch from torch.utils.data import TensorDataset def load_mnist_images(filename): with gzip.open(filename, 'rb') as f: magic = int.from_bytes(f.read(4), 'big') num_images = int.from_bytes(f.read(4), 'big') rows = int.from_bytes(f.read(4), 'big') cols = int.from_bytes(f.read(4), 'big') data = np.frombuffer(f.read(), dtype=np.uint8) return data.reshape(num_images, rows, cols) def load_mnist_labels(filename): with gzip.open(filename, 'rb') as f: magic = int.from_bytes(f.read(4), 'big') num_labels = int.from_bytes(f.read(4), 'big') data = np.frombuffer(f.read(), dtype=np.uint8) return data def build_dataset(image_path, label_path): images = load_mnist_images(image_path) labels = load_mnist_labels(label_path) images_tensor = torch.tensor(images, dtype=torch.float32).unsqueeze(1) / 255.0 labels_tensor = torch.tensor(labels, dtype=torch.long) return TensorDataset(images_tensor, labels_tensor)

这段代码用gzip和numpy.frombuffer直接解析 IDX 格式,不做 PIL 转换、不依赖 torchvision 的下载接口。unsqueeze(1)在通道维上增加一维,使数据形状从(N, 28, 28)变成(N, 1, 28, 28),匹配 PyTorch 卷积层的输入要求。像素值除以 255.0 归一化到 0~1 区间,这是影响收敛速度的关键细节——不归一化直接输入的话,梯度容易震荡。TensorDataset将图片和标签打包,供 DataLoader 迭代使用。

3. 把 FedAvg 核心算法拆解到参数级:服务端聚合与客户端更新的联动机制

3.1 客户端更新的三要素:本地 epoch、batch size 与学习率

本地训练的质量取决于三个参数:本地 epoch 数、batch size、学习率。代码里通过在clients.py的local_train函数中调用model.state_dict()来获取更新后的权重字典。这里有一个容易被忽略的细节,state_dict()返回的是参数的浅拷贝引用,客户端在返回前需要确保已经断开了梯度计算,否则序列化传输时会报错。

从代码可以看出,本地 epoch 数直接控制客户端计算开销。epochs=5表示每个客户端在自己的私有数据上完整迭代 5 轮,这种方式相比每轮只做一次梯度更新,能显著减少服务端与客户端之间的通信次数。但注意 local epoch 过大会导致「本地过拟合」——每个客户端在自己的数据分布上收敛过头,聚合后反而损害全局性能,这在 non-iid 数据下尤其明显。

3.2 服务端聚合的加权逻辑与全局模型更新公式

服务端的聚合逻辑集中在server.py的核心循环中。FedAvg 的本质是对各个客户端返回的模型参数做加权平均,权重是每个客户端持有的样本量占总样本量的比例。代码中用如下方式实现:

def fedavg_aggregate(global_model, client_updates, client_sizes): total_size = sum(client_sizes) global_dict = global_model.state_dict() for key in global_dict.keys(): weighted_sum = 0.0 for client_state, client_size in zip(client_updates, client_sizes): weighted_sum += client_state[key].float() * (client_size / total_size) global_dict[key] = weighted_sum global_model.load_state_dict(global_dict) return global_model

这段代码的key遍历了网络中的所有参数层,包括卷积核权重和偏置项。(client_size / total_size)是聚合权重,样本多的客户端在全局模型中拥有更大的发言权。这里要注意加权平均过程中的数字精度:PyTorch 的默认张量类型是 float32,在累加多个客户端更新时可能出现精度损失,特别是在模型接近收敛后更新量很小的情况。一个常见做法是先求和再除以总样本数,而不是逐项计算比例后累加。

3.3 参数配置对照表与通信轮次的设置建议

运行联邦训练前,需要理清几个关键参数的推荐范围。以下是常用配置:

参数推荐范围对训练的影响
总客户端数10~100太大时单轮通信成本上升
每轮参与比例0.1~0.5比例过低导致聚合不稳定
本地 epoch 数1~10过大易局部过拟合
batch size16~64影响本地 SGD 收敛质量
全局轮次50~200视模型收敛曲线而定
学习率0.01~0.1过高震荡,过低收敛慢

这个项目里总客户端数由dataSets.py中的切分数决定,每轮参与比例在server.py中控制。如果你要模拟大规模联邦场景,可以把客户端数增大,但要同步考虑内存占用——每个客户端保留一份完整模型权重,100 个客户端就是 100 份拷贝,对内存不太友好。

3.4 聚合时机与异步联邦的边界

部分读者在跑通同步 FedAvg 后会追问「能不能改造成异步」。当前项目的实现是同步聚合——服务端必须等本轮所有参与客户端返回参数后才能继续下一轮。这在真实场景中会遇到「掉线客户端」问题,但作为研究用代码,同步假设是可以接受的简化。如果你需要异步联邦,需要额外处理过期更新和延迟容忍机制,当前这份代码没有涉及。FedAvg 的代码价值在于它展示了最基本的联邦学习闭环,异步化改造需要你自己扩展。

4. 数据切分的门道与 non-iid 实现的背后逻辑

4.1 训练集测试集的拆分策略与代码映射

MNIST 原始数据集由 60000 张训练图片和 10000 张测试图片组成。联邦学习场景下,这 60000 张训练图需要被分配到不同的客户端手里,分配方式直接决定实验是 iid 还是 non-iid。

def partition_data(dataset, num_clients, noniid_ratio=0.5): num_samples = len(dataset) indices = list(range(num_samples)) client_data = [[] for _ in range(num_clients)] num_noniid = int(num_clients * noniid_ratio) num_iid = num_clients - num_noniid # non-iid 客户端:按标签排序后分段分配 sorted_by_label = sorted(indices, key=lambda i: dataset[i][1]) samples_per_client = num_samples // num_clients for c in range(num_noniid): start = c * samples_per_client end = start + samples_per_client client_data[c] = sorted_by_label[start:end] # iid 客户端:随机打乱后均匀分配 remaining = indices[num_noniid * samples_per_client:] random.shuffle(remaining) for idx, sample_idx in enumerate(remaining): c = num_noniid + (idx % num_iid) client_data[c].append(sample_idx) return client_data

这里的noniid_ratio=0.5表示一半客户端持有了标签分布严重倾斜的数据,另一半客户端持有了均匀分布的数据。代码先按标签排序再分段,本质上是把某些客户端的数据限制在少数几个数字类别中。比如编号 0 的客户端可能只持有数字 0、1、2 的样本,而编号 1 的客户端可能只持有数字 3、4、5 的样本。这就是 non-iid 的核心含义。

4.2 标签分布倾斜的影响与后续模型性能变化

non-iid 划分直接影响聚合模型的质量。一个只见过数字 0 的客户端,它在本地训练时会把所有参数推向「偏向数字 0」的方向;服务端把这些偏斜的参数平均后,全局模型可能在某些类别上表现好,在另一些类别上表现差。这种现象最早在联邦学习研究中被称为「客户端漂移」,本质与持续学习中的灾难性遗忘类似。

观察方式也很直接——每个客户端持有标签类别的直方图,如果直方图接近均匀分布,那就是 iid;如果极不均衡,就是 non-iid。建议在实验记录中保留一张客户端标签分布表,方便复现时对照。

4.3 数据划分随机性带来的复现一致性问题

复现实验时最容易被忽视的问题是:random.shuffle不设种子会导致每次运行的数据划分不同。同样是 non-iid 实验,昨天跑出的性能和今天跑出的性能不可比。解决方法是设置全局随机种子:

import random import numpy as np import torch def set_seed(seed=42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)

这个函数建议在server.py的入口处调用。设置固定随机种子后,数据划分、模型初始化、客户端采样顺序全部确定,不同轮次的实验结果可以互相比较。这是做联邦学习实验的一条基本纪律。

5. 避坑排查:MNIST 联邦学习实战中常见的六个翻车现场

5.1 MNIST 数据下载 404 与 torchvision 接口不稳定的问题

现象:运行python getData.py或直接用torchvision.datasets.MNIST下载数据时,报 HTTP 404 错误,或者下载到一半中断。

原因:MNIST 官方源在部分网络环境下不可达,torchvision 内置的下载链接经常失效,这是社区里反复出现的问题。

解决:不使用自动下载,手动下载四个 gz 文件放到data/目录下,再用dataSets.py直接读取本地文件。如果只有压缩包内的数据,检查文件名是否完全对应:t10k-images-idx3-ubyte.gz、t10k-labels-idx1-ubyte.gz、train-images-idx3-ubyte.gz、train-labels-idx1-ubyte.gz。注意不要解压后改名,dataSets.py依赖 gz 二进制流解析。

5.2 PyTorch 版本差异导致state_dict加载失败

现象:客户端返回的state_dict在服务端load_state_dict时提示size mismatch或missing keys。

原因:客户端模型和服务端模型结构定义不一致,或者 PyTorch 版本间存在参数命名差异。更隐蔽的情况是,某些层使用了不同的初始化方式导致张量形状不同。

解决:在server.py初始化全局模型后,建议打印第一层卷积的参数形状,再与客户端Models.py中定义的形状对比。排查顺序是:模型类是否同一个类 → 是否有额外的Dropout或BatchNorm层 → 是否在客户端加载过预训练权重。

5.3 本地 epoch 过大导致全局模型性能暴跌

现象:本地 epoch 从 5 改到 20 后,全局模型的测试准确率反而下降了 15~20 个百分点。

原因:每个客户端在自己的本地数据上训练太久,模型参数向客户端各自的数据分布偏移过多,聚合时产生了互相抵消的更新。这在 non-iid 分布下尤其严重,相当于每个客户端都在追逐自己的局部最优解,全局模型被拉向你推我挤的混乱状态。

解决:把本地 epoch 控制在 1~3 之间。如果业务场景要求更多本地迭代,需要引入本地正则化项,比如在损失函数中加入全局模型与本地模型的 KL 散度约束,但这份代码没有实现该机制,需要自行扩展。

5.4 客户端数量过多导致的内存溢出

现象:客户端数从 10 加到 100 后,MemoryError直接中断训练,内存占用飙升到几个 GB。

原因:每个客户端持有一份独立的模型参数副本,clients.py中每个client_state都是一个完整的state_dict。100 个客户端就是 100 份权重副本,再加上梯度计算的开销,内存吃不消。这个问题可以通过参数总量简单计算——一个单层卷积网络加上全连接层大约 2~5 万参数,float32 存储每个参数 4 字节,100 个客户端总共约 20MB,看起来不大,但实际操作中 PyTorch 的自动微分图占用了额外内存,实际开销远大于理论值。

解决:控制每轮参与比例,不要让所有客户端同时返回状态字典。也可以研究服务端聚合的流式处理方式,收到一个客户端就聚合一次,而不是全部接收后再统一聚合。

5.5 测试阶段错误使用了客户端本地数据

现象:在server.py中做全局模型评估时,直接用了某个客户端的数据,导致准确率评估结果飘忽不定,不同客户端上差异极大。

原因:联邦学习的测试集应该与所有客户端训练数据严格隔离。如果用了某个客户端的私有数据做评估,那就是在「留出法」上开了个口子——全局模型可能过拟合了该客户端的分布。

解决:从 MNIST 原始测试集中取 10000 张作为独立测试集,这部分不参与任何客户端的数据分配。在联邦学习论文中常用「全局测试集」概念,指的就是这部分不落入任何客户端的数据。

5.6 加权平均时张量类型不匹配的隐性报错

现象:聚合代码运行时偶发RuntimeError: Expected object of scalar type Float but got scalar type Double,而且不是固定的轮次出现。

原因:部分客户端返回的模型参数是 float64 类型,而全局模型是 float32。通常这是因为某个客户端本地数据进行了高精度转换,或者不同客户端使用了不同的 PyTorch 默认类型设置。

解决:在聚合前统一类型转换,或者干脆在local_train返回前强制.float()。建议在fedavg_aggregate函数里对所有键值做client_state[key] = client_state[key].float(),避免类型问题在训练中途随机爆发。

6. 收敛效果自检:一种不用外部框架就能完成的联邦模型验证法

联邦学习跑完之后,怎么证明全局模型真的学到了知识?最简单的指标是整体准确率,但这个数字掩盖了分客户端、分标签类的性能差异。我的习惯做法是在server.py末尾追加一个细粒度的评估函数——只统计模型在「从未参与训练的测试集」上的表现,并且按数字 0~9 分开统计。

def evaluate_per_class(model, test_loader, device): model.eval() correct = [0] * 10 total = [0] * 10 with torch.no_grad(): for data, target in test_loader: data, target = data.to(device), target.to(device) output = model(data) pred = output.argmax(dim=1) for i in range(10): mask = target == i total[i] += mask.sum().item() correct[i] += (pred[mask] == target[mask]).sum().item() return {i: correct[i] / total[i] for i in range(10) if total[i] > 0}

这个函数执行按类别的正确率统计。如果某些类别准确率明显低于平均值,说明某个客户端的 non-iid 数据分布主导了该类别训练,全局模型在该类上存在系统性偏差。对比不同客户端数量配置下的 per-class 准确率,还能看到数据分布对模型公平性的影响。

过了自检这一步,还可以做一次更严格的验证:把全局模型的参数作为初始化权重,在完整 MNIST 训练集上微调一个 epoch,对比微调前后的准确率提升速度。如果微调初始阶段收敛显著快于随机初始化,说明全局模型已经携带了有效的特征提取能力。这个技巧的成本极低,但能直观验证联邦聚合的收敛质量。

我在复现这个项目时养成的习惯是:每跑完一组配置,先在测试集上输出 per-class 准确率矩阵,再决定要不要调整数据切分方式。从那以后我每次准备联邦学习实验都会强制走一遍这个流程——确认分类均衡性、检查类型一致性、记录随机种子。希望这套排查思路帮你在复制 FedAvg 工程时少走几趟弯路。

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

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

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

立即咨询