☰
横向联邦图像分类实战:PyTorch实现FedAvg与CIFAR10数据划分
2026/9/28 13:15:49 网站建设 项目流程

简介:这是一份基于Python从零实现横向联邦图像分类的完整学习代码,适配联邦学习实战图书第3章内容,适合正在入门联邦学习、需要动手实践图像分类场景的在校学生、研究者和开发人员。代码模块清晰,涵盖数据集加载、模型构建、客户端与服务器协同训练等核心步骤,并配有大量注释,方便逐行理解横向联邦的工作流程。资源共22个文件,以6个Python源码文件为主,辅以配置文件、编译缓存、说明文档和示意图,压缩包仅156KB,轻量易部署。内部包含README说明、JSON配置和网络结构图,可帮助快速跑通实验并对照理解参数设置。全部代码已测试通过,答辩评审平均分达96分,可直接作为课程作业或项目起步参考。目前已有125人学习浏览,对于想快速掌握横向联邦图像分类实现细节的读者,这份代码能节省搭建环境与调试流程的时间,尤其在客户端-服务器通信与模型聚合方面能提供直接借鉴。

1. 横向联邦图像分类的落地样板:这份代码练的是哪三件事

如果你手里已经有一套能跑通的 PyTorch 图像分类代码,想把它改造成横向联邦版本,你会发现真正的难点根本不在模型,而在数据怎么分、参数怎么传、多个客户端怎么协作、服务端怎么聚合。这套代码恰好把这几件事拆成了可以逐行阅读的模块:入口是main.py,配置集中在utils/conf.json,服务端在server.py,客户端在client.py,模型和数据集划分单独放在models.py与datasets.py。它对应联邦学习实战书中第三章的横向联邦图像分类章节,全程用 Python + PyTorch 从零实现,训练数据是 CIFAR10,跑通之后可以直观看到每一轮全局模型准确率的抬升。适合两类人:刚接触联邦学习、想找一个最小可运行工程的学生,以及想把自己的图像分类任务套上横向联邦框架、需要动数据划分逻辑的工程师。

2. 环境与数据先行:两小时内把第一个联邦轮次跑起来

拿到源码包第一步不是看代码,是先把环境固下来。横向联邦的代码结构比普通单机训练多了一层“服务端—客户端”的抽象,如果 Python、PyTorch、数据集三者的版本不匹配,后面排查起来会非常痛苦。我拆过的这种带大量注释的教学代码,作者基本都在 Python 3.8 上开发测试,PyTorch 版本则集中在 1.10 到 1.13 之间,所以建议你不要一上来就装最新的 Python 3.12 或 PyTorch 2.x,能跑,但没必要给自己增加变量。

2.1 环境准备:Python、PyTorch、torchvision 的版本搭配

先建一个独立的 conda 环境,避免把系统 Python 搞乱。这个习惯在跑任何带工程结构的项目里都值得保留,尤其是联邦学习这种要反复改配置、反复起实验的场景:

conda create -n fl python=3.8 -y conda activate fl pip install torch==1.13.1 torchvision==0.14.1 pip install numpy

这段命令做了三件事:创建名为fl的 Python 3.8 环境,安装 CPU/GPU 通用的 PyTorch 1.13.1 与配套 torchvision,再装一个 numpy 兜底。之所以锁死 torchvision 版本,是因为 torchvision 和 torch 的编译版本必须配套,torchvision.datasets.CIFAR10这个接口在不同版本里行为基本一致,但 pyc 缓存和算子派发偶尔会闹脾气——你在自己的环境里用 2.x 跑也不是不行,只是后面遇到算子报错时,先回头确认 torch 和 torchvision 是否配套,永远是第一步。

需要留意的是,如果你机器上有多个 CUDA 版本,pip install torch默认装的可能是 CPU 版或与你系统 CUDA 不匹配的版本。检查方式很简单:启动 Python 后执行import torch; print(torch.cuda.is_available()),输出False说明当前装的是 CPU 版,CIFAR10 这种小图用 CPU 硬跑几十轮也能出结果,但多客户端轮流训练会非常慢,建议还是装对应 CUDA 版本的 PyTorch。

2.2 CIFAR10 数据放置:直接放 data 文件夹还是自动下载

这份代码在datasets.py里已经内置了 CIFAR10 下载逻辑,但你最好先把数据准备好,再跑主程序,否则第一轮训练会卡在下载环节,还要处理断点续传的问题。

from torchvision import datasets, transforms transform = transforms.Compose([transforms.ToTensor()]) train_set = datasets.CIFAR10(root='./data', train=True, download=True, transform=transform) test_set = datasets.CIFAR10(root='./data', train=False, download=True, transform=transform)

这里root='./data'的含义是数据会被放到当前工作目录的data文件夹下。很多人跑出FileNotFoundError不是代码问题,而是当前终端目录不在项目根目录,相对路径自然就指错了。稳妥的做法是把这个root改成项目根目录的绝对路径,或者在启动main.py之前先用cd进到项目目录。CIFAR10 整个数据集解压后大概是 170M 左右,训练集 50000 张图片,测试集 10000 张,下载源在国外,网络差的时候容易卡在半路。如果下载总是失败,去 CIFAR10 官网手动下载cifar-10-python.tar.gz,解压后把cifar-10-batches-py文件夹整个放进data目录,代码照样能读。

2.3 目录结构:每个文件的职责在动手前先看清楚

把源码包解开之后,我一般会按main.py -> server.py -> client.py -> models.py -> datasets.py这个顺序读一遍,而不是从README.md开始。原因很简单:README往往只讲怎么运行,不讲数据流;但代码文件之间的调用关系才是理解联邦学习项目的关键。下面是这份资源里最容易混淆的几个文件的职责划分:

文件职责重要程度
main.py解析命令行配置、组织多轮联邦训练流程入口核心
server.py初始化全局模型、聚合客户端上传的权重必须理解
client.py加载本地数据、执行本地训练、返回权重增量必须理解
models.py定义图像分类网络(LeNet5 等)按需修改
datasets.py把 CIFAR10 按 IID / Non-IID 划分给各个客户端实验重点
utils/conf.json超参、路径、模型选择等集中配置每次实验必改

运行命令非常简短:

python main.py -c ./utils/conf.json

-c参数指定配置文件路径,main.py内部用 argparse 读取 conf.json 并把里面的参数分发给服务端和客户端。第一次运行如果数据没下载过,终端会先打印一堆 CIFAR10 的下载进度条,然后是数据划分日志,接着才会出现训练进度。看到这里,你就可以确认整个链路已经通了,下一步就是拆解每一层在干什么。

3. 看懂启动链路:main.py 的入口、conf.json 的参数与服务端/客户端职责

横向联邦和单机训练最大的区别在于:单机训练的循环是“前向—反向—更新”三层,而联邦学习的循环是“客户端本地训练—上传—服务端聚合—广播”四步。main.py 的作用不是训练模型,而是把这几步按轮次串起来。我见过不少读者卡在这一层,觉得 main.py 怎么没有 loss.backward(),其实反向传播全部在 client.py 里,main.py 只负责调度。

3.1 conf.json:横向联邦的超参集中营

打开utils/conf.json,你会看到一份类似下面的配置。字段名可能略有出入,但横向联邦实验绕不开的就是这几个维度的参数:

{ "data_path": "./data", "num_clients": 10, "frac": 0.3, "rounds": 50, "local_epochs": 2, "batch_size": 64, "lr": 0.01, "model": "LeNet5", "device": "cuda" }

参数的选择直接决定实验行为:num_clients是模拟的客户端总数,frac是每轮参与训练的客户端比例,rounds是联邦轮次总数,local_epochs是每个客户端本地训练几轮,batch_size和lr与单机训练相同,model指定使用哪个网络结构。这些参数之间的配合关系比参数本身更重要:frac=0.3 + num_clients=10意味着每一轮只有 3 个客户端被随机选中参与聚合,客户端数量越少、local_epochs越大,全局模型的收敛就越不稳定。初学者最容易犯的错误是把local_epochs当成普通训练轮次直接调到 10 以上,结果全局模型在客户端之间“漂移”,准确率来回震荡。

3.2 argparse 解析:-c ./utils/conf.json是怎么生效的

main.py 的入口参数解析逻辑一般长这样,这段代码同时也是 Python 命令行工具的标准写法:

import argparse import json parser = argparse.ArgumentParser(description='横向联邦图像分类') parser.add_argument('-c', '--config', required=True, type=str, help='配置文件路径,例如 ./utils/conf.json') args = parser.parse_args() with open(args.config, 'r', encoding='utf-8') as f: conf = json.load(f)

这段代码的逻辑很直白:argparse定义了一个-c参数,用户在命令行传入的路径被保存到args.config,随后通过open + json.load把字典读进conf变量。required=True表示如果不传-c参数,程序会直接报错退出,这种设计在工程上是故意“早失败”,避免后面代码里出现一长串None导致的隐性 bug。有一个细节很多人会踩:如果 conf.json 里写了中文注释,open时必须带上encoding='utf-8',否则在 Windows 上会默认用 GBK 解析,直接抛UnicodeDecodeError。

3.3 训练主流程:每个 round 里到底发生了什么

读完配置,再看 main.py 的调度逻辑。横向联邦的标准流程可以用这个伪代码描述,这份资源的 main.py 结构我猜测大概率也是这个骨架,只是函数名做了拆分:

import random global_model = server.init_model(conf) for round_id in range(conf["rounds"]): # 按比例挑选本轮参与客户端 num_selected = max(1, int(conf["num_clients"] * conf["frac"])) selected = random.sample(range(conf["num_clients"]), num_selected) updates = [] for cid in selected: # 客户端基于全局模型做本地训练,返回权重更新 local_update = client.local_train(cid, global_model, conf) updates.append(local_update) # 服务端聚合所有客户端上传的更新 global_model = server.aggregate(global_model, updates) print(f"round {round_id + 1}/{conf['rounds']} finished")

这段代码里最关键的一行是client.local_train的返回值。它返回的不是完整模型,而是“权重增量”,也就是本地训练后的参数减去初始全局参数之间的差值。服务端拿到多个客户端的增量后,做加权平均再更新全局模型,这就是 FedAvg 的基本思想。理解这一点对后面改代码很有帮助:如果你想在聚合时给不同客户端加上不同权重,改动点就在server.aggregate的加权逻辑里;如果你想改变客户端选择策略,改动点就在random.sample这行。整个横向联邦的黑匣子,拆开之后就是“一个循环 + 一次平均”。

4. 核心模块拆解:models.py 的网络结构、datasets.py 的数据划分与服务端聚合

前两章讲的是“能跑”,这一章讲的是“能改”。很多读者下载这份代码是为了做毕设或课程设计,如果把 models.py 和 datasets.py 看懂,你就能把 CIFAR10 换成自己的数据集,把 LeNet5 换成 ResNet,把 IID 划分改成 Non-IID 划分——这才是这份资源真正的价值点。

4.1 models.py:图像分类网络为什么这样设计

这套代码里最基础的分类型号是 LeNet5,一个在 CIFAR10 这种小图上足够用的卷积网络。CIFAR10 每张图是 3 通道 32×32,而 LeNet5 原始设计输人就是 32×32 灰度图,所以这里只需要把第一层卷积的输入通道从 1 改到 3,其余结构基本可以保留:

import torch.nn as nn import torch.nn.functional as F class LeNet5(nn.Module): def __init__(self, num_classes=10): super().__init__() self.conv1 = nn.Conv2d(3, 6, 5, padding=2) # 3 通道输入,32x32 保持尺寸 self.conv2 = nn.Conv2d(6, 16, 5) # 16 通道,尺寸自动变小 self.fc1 = nn.Linear(16 * 6 * 6, 120) self.fc2 = nn.Linear(120, 84) self.fc3 = nn.Linear(84, num_classes) 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)) x = F.relu(self.fc2(x)) return self.fc3(x)

注意conv1的padding=2,这是为了让 32×32 的输入经过 5×5 卷积后仍然保持 32×32,再配合 2×2 最大池化降维到 16×16;第二层卷积没有 padding,16×16 经过 5×5 卷积后变成 12×12,再池化变成 6×6。因此全连接层第一维是16 * 6 * 6,这个数字是亲手算出来的,不是随便拍的。如果你想换成models.py里可能提供的 ResNet18,有两个地方必须同步改:第一个是输入层,ResNet18 默认接受 224×224 的输入,CIFAR10 的 32×32 要么在数据增强里 resize,要么把第一层卷积的 stride 和 kernel 改掉;第二个是最后的全连接层,把num_classes改成你数据集的类别数。

4.2 datasets.py:IID 与 Non-IID 数据划分的差异

横向联邦和普通图像分类在数据上的本质区别是“数据不在同一个地方”。datasets.py 解决的就是这个问题:把 CIFAR10 的 50000 张训练图按一定规则切分给 10 个客户端。最简单的划分方式是 IID,也就是把所有样本打乱后均匀分配,每个客户端手里的数据分布都接近全局分布。这种做法在实现上最简单,但和真实联邦场景差距很大——现实里各客户端的数据往往是偏态的。

import random def split_iid(dataset, num_clients): idxs = list(range(len(dataset))) random.shuffle(idxs) step = len(dataset) // num_clients return [idxs[i * step:(i + 1) * step] for i in range(num_clients)]

这段代码先把所有样本索引打乱,再按客户端数量平均切成若干块,每个客户端拿到数量相等、分布近似的数据。如果你只做毕设演示,IID 划分就能得到一个还不错的收敛曲线。但如果你想模拟真实场景,datasets.py 里大概率还提供了 Non-IID 划分,常见做法是按标签排序后使用狄利克雷分布(Dirichlet)分配,让每个客户端只拥有少数几个类别的样本:

import numpy as np from collections import defaultdict def split_dirichlet(dataset, num_clients, alpha=0.5): label_to_idx = defaultdict(list) for idx, (_, label) in enumerate(dataset): label_to_idx[label].append(idx) # 每个客户端按 Dirichlet 分布抽样每种标签的样本 client_idx = [[] for _ in range(num_clients)] for label, idxs in label_to_idx.items(): np.random.shuffle(idxs) proportions = np.random.dirichlet([alpha] * num_clients) # 按比例切分当前标签的数据 assigned = 0 for cid in range(num_clients): take = int(len(idxs) * proportions[cid]) client_idx[cid].extend(idxs[assigned:assigned + take]) assigned += take return client_idx

alpha是 Non-IID 程度的控制旋钮:alpha 越大,每个客户端的标签分布越接近全局;alpha 越小,分布越偏,训练难度也越高。我在做对比实验时习惯固定住随机种子,然后用 alpha 从 0.1 到 1.0 做一组扫描,观察全局模型的收敛速度变化。这个实验做完,对横向联邦“数据异构导致模型漂移”这句话的理解会非常直观。

4.3 server.py 中的聚合逻辑:联邦平均到底平均了什么

最后一块核心是服务端聚合。很多第一次接触联邦学习的人会以为聚合是把多个模型的参数直接求平均,实际上这里有个比参数平均更微妙的细节:应该平均完整参数,还是平均权重增量?两种做法都能收敛,但增量聚合更稳定,因为它天然剔除了全局模型本身的“底数”,服务端只需要把这些差值叠加回当前全局模型即可。

import torch def aggregate(server_model, client_weights): # client_weights 是多个客户端返回的 state_dict 列表 global_dict = server_model.state_dict() for name in global_dict.keys(): global_dict[name] = torch.mean( torch.stack([w[name].float() for w in client_weights]), dim=0 ) server_model.load_state_dict(global_dict) return server_model

这段代码的逻辑是:遍历全局模型的每一个参数张量,把本轮参与训练的所有客户端对应参数堆叠起来,在维度 0 上求平均,最后整体载回服务端模型。torch.stack要求所有客户端的state_dict拥有完全相同的键和形状,所以客户端在本地训练之前,必须把服务端广播的全局模型原样加载进去,不能自己加层、减层。另一个容易被忽略的问题:如果用动量 SGD 作为本地优化器,每个客户端的optimizer.state里保存的动量缓冲并不会上传到服务端,下一轮开局时客户端是用全新的优化器状态去继承上一轮聚合出的全局模型。这意味着local_epochs不宜设置太大,否则本地训练方向会越偏越远,回到全局模型时反而破坏整体精度。

5. 复现中的避坑清单:四个高频报错与排查路线

所有联邦学习项目第一次跑都会出问题,这不是代码质量差,而是多了一个“服务端—客户端”维度之后,普通的单机训练经验有一半不适用了。我在复现这类项目时踩过不少坑,下面这几条是最常见的,逐条按「现象 → 原因 → 解决」说清楚,你可以当作排查手册用。

5.1 坑一:FileNotFoundError: No such file or directory: 'data/cifar-10-batches-py/data_batch_1'

现象:执行python main.py -c ./utils/conf.json后,程序没有进入训练,直接抛出文件不存在错误。

原因:这个报错 80% 不是因为数据没下载,而是当前工作目录不在项目根目录。你在 PyCharm 或 VS Code 里直接点运行按钮时,默认的工作目录可能是项目根目录,也可能是配置文件的所在目录,两种情况下相对路径./data指向的位置完全不同。

解决:先在终端用cd进入项目根目录,再执行运行命令;同时把utils/conf.json里的data_path改成绝对路径。我个人的习惯是直接改成项目根目录拼接:"data_path": "/Users/xxx/wqf-federal-learning-master/data",一劳永逸,不再受工作目录影响。

5.2 坑二:RuntimeError: torchvision.datasets.CIFAR10下载到一半卡住,或者解压时报错

现象:第一次运行时下载进度条走了一半不动,或者下载完成后报EOFError/ 压缩包损坏。

原因:CIFAR10 的下载源在国外服务器,大文件下载经常中断,download=True不会做断点校验,数据不完整就会在解压时报错。

解决:断开后先把data目录下残留的临时文件删干净,再手动用浏览器下载cifar-10-python.tar.gz,解压出cifar-10-batches-py文件夹放到data下。这样 torchvision 检测到本地已有完整数据,会跳过下载直接读取。项目里还提到了figures目录里的fig31.png这类图表,那是作者实验曲线,不影响运行,不用额外处理。

5.3 坑三:客户端本地准确率挺高,但服务端全局模型准确率上不去

现象:每个客户端在本地测试集上的 acc 能到 85% 以上,但跑完 50 轮联邦后,服务端全局模型在测试集上只有 60% 左右,甚至更低。

原因:这是横向联邦里最典型的“客户端漂移”现象。原因往往不是聚合代码写错了,而是local_epochs太大或学习率太高,导致每个客户端在本地把模型推向了只适合自己数据的区域;服务端把这几个方向各异的模型一平均,自然互相抵消。另一个常见原因是每轮参与客户端太少,随机波动太大。

解决:先把local_epochs调回 1,lr降低到 0.001~0.01,frac提高到 0.5 以上,看全局准确率是否稳定上升。如果还是不行,在server.aggregate里打印参与聚合的客户端数量,确认每轮真的有多个客户端返回,而不是只选到了 1 个客户端。

5.4 坑四:多客户端同一块 GPU,跑到第 N 轮突然 OOM 显存溢出

现象:前几轮一切正常,训练到中途报CUDA out of memory,程序退出后显存占用却迟迟不释放。

原因:device直接写死成cuda,所有客户端共享同一个 GPU,但本地训练时保存了多个模型的梯度图、优化器状态和验证数据,这些对象不会自动释放。第 N 轮溢出通常不是突然变大,而是前几轮的临时变量一直没清干净,累积到临界点才爆掉。

解决:如果代码是串行训练客户端,每轮客户端训练结束后,显式调用torch.cuda.empty_cache(),并把客户端模型del掉或搬回 CPU;如果代码是并行训练多个客户端,最简单的方法是限制num_clients为实际可用的显卡数量,batch_size同步降低到 32 甚至 16。做实验阶段不需要多客户端并行,串行跑反而稳定,慢但不会 OOM。

5.5 坑五:argparse报unrecognized arguments或the following arguments are required: -c

现象:启动命令照抄 README,但报参数解析错误;或者自己加了--model resnet18直接报 unrecognized。

原因:main.py里只定义了-c / --config一个参数,其他所有参数都通过 conf.json 传入。直接把 conf.json 里的字段挂到命令行参数上,argparse 当然不认。

解决:不要在命令行加额外参数,要改模型或超参就改utils/conf.json;如果不习惯 JSON 写法,可以在main.py的argparse里额外定义几个可选参数,再用args覆盖 conf 里对应字段。最常见的做法是只保留-c,全部通过配置文件驱动,这样每次实验的记录也更干净。

6. 从复现到改造:打印每一轮日志、切换模型、做你自己的实验

代码能原样跑通只是开始。我建议你做的第一件事不是改逻辑,而是在main.py里加一行验证日志,把每一轮的关键指标打印出来,方便后续对比实验。

acc = evaluate(global_model, test_loader) print(f"round {round_id + 1}/{conf['rounds']}, " f"global_acc={acc:.4f}, " f"active_clients={len(selected)}")

我在复试到main.py的时候,一定会把这个日志写到文件里而不是只看终端,因为联邦实验的轮次一旦超过 50,终端滚屏记录很难回溯。用tee命令或者logging把输出重定向到train_log.txt,之后画准确率曲线就是一行matplotlib的事。

第二个值得做的实验是“异构程度扫描”。在datasets.py里把split_dirichlet的alpha分别设为 0.1、0.5、1.0,跑完全相同的轮次,对比全局模型在测试集上的最终精度。你会发现 alpha 越小,收敛越慢,最终精度也有明显下降——这就是横向联邦在真实场景下面临的核心挑战。做这个实验时,记得固定随机种子,并且每次只改 alpha 一个变量,否则对比结果不干净。

模型层面,如果你想换models.py里的网络,注意两点:CIFAR10 的 32×32 输入对 ResNet 类网络并不友好,常见的做法是在数据增强里加RandomCrop(32, padding=4)和RandomHorizontalFlip();另外全连接层的输出维度要改成num_classes。联邦平均对模型结构没有特殊要求,只要所有客户端都用同一个网络结构,服务端就能聚合。基于这份源码,你能做的扩展还很多:把 CIFAR10 换成你自己的图片分类数据集,只需在datasets.py里替换数据加载逻辑并保持索引划分结构;把随机选择客户端改成按数据量加权选择;把聚合方式从简单平均改成 FedProx 的正则化版本。从那以后我每次跑联邦学习项目,都强制走一遍固定版本、固定数据路径、固定随机种子三个动作,之后才敢动训练参数。看似多花十分钟,实际上省掉的是一整晚的排障时间。希望帮到你。

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

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

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

立即咨询