简介:代码基于Python从零实现横向联邦图像分类,是《联邦学习实战》第3章的配套学习工程,面向联邦学习初学者及计算机相关专业学生,适合有一定Python基础、希望深入理解联邦学习实战细节的读者。代码注释极为详尽,涉及数据加载、模型构建、客户端与服务端通信及参数聚合等核心逻辑,运行仅需PyTorch和CIFAR-10数据集,配置参数通过JSON文件统一管理,并附有清晰的运行入口、配置说明与指导。压缩包共22个文件,以6个Python脚本为主体,辅以JSON配置、Markdown说明文档、结果图片、项目配置文件及运行生成的缓存文件,整体仅156KB,轻量便携且目录结构清晰,便于按模块研读。目前已有126人学习浏览。代码经测试运行成功,完整演示了横向联邦图像分类从数据分发、本地训练到模型聚合的全流程,且未依赖高层联邦框架,手动实现核心交互逻辑,是理解联邦学习底层机制的理想范例;亦可作为课程设计或毕业设计的参考工程,作者提供远程教学与答疑支持,遇到环境搭建或原理疑问时可获得针对性帮助。
1. 横向联邦图像分类:这份Python学习代码能帮你跑通的第一关
横向联邦图像分类是联邦学习入门绕不开的一个实战点:多个客户端各自持有本地图片数据,不下传原始样本,只通过服务端聚合模型参数来共同训练一个图像分类模型。这份资源就是《联邦学习实战》第3章的配套代码,用PyTorch从零手写了完整的横向联邦图像分类流程,附带大量中文注释。不少初学者拿到源码包后卡在第一步——不知道先装什么、数据集放哪、为什么main.py报路径错误。我把这套代码按“环境准备 → 代码结构 → 跑通训练 → 避坑 → 进阶调试”的顺序拆开讲,重点落在每个文件的实际作用和能直接照做的命令上。适合刚学完深度学习基础、想理解参数聚合过程的在校生,也适合做课程设计或毕设演示的从业者。建议动手前先把README.md和utils/conf.json各看一遍,能少踩一半的坑。
2. 环境准备与数据落盘:从Python安装到CIFAR-10就位
2.1 Python、PyTorch环境搭建与版本选型
这套代码基于Python 3和PyTorch实现。作者在摘要里要求先装好Python、PyTorch环境,再下载CIFAR-10数据集。我习惯用conda先创建一个干净的虚拟环境,避免多个项目之间的依赖打架。Python版本选3.8比较省心,PyTorch 2.0系列的API与这份代码的用法基本兼容,如果你机器上已经有更旧的1.x版本,多数情况下也能跑,只是部分接口名可能有差异。
conda create -n fed python=3.8 -y conda activate fed pip install torch torchvision创建虚拟环境后,torch和torchvision会一起装上。torchvision里自带了CIFAR-10数据集的下载与读取工具,这份代码的datasets.py大概率就是基于torchvision.datasets.CIFAR10来做加载的。如果机器有NVIDIA显卡,建议装对应CUDA版本的PyTorch,没有的话CPU版本也可以运行——只是训练速度会慢一些,对理解联邦学习流程没有影响。
装完环境后,可以先跑一句简单的Python命令确认PyTorch能正常导入:
python -c "import torch; print(torch.__version__)"能打印出版本号,说明基础环境没问题。接着就需要准备数据了。
2.2 CIFAR-10数据集下载、目录结构与校验
代码里读取数据时找的是data文件夹。你需要在项目根目录下建一个data目录,把CIFAR-10数据集放进去。最常见的做法是直接用torchvision下载并解压,命令在Python交互环境或脚本里执行都行:
import torchvision # 下载到项目根目录下的 data 文件夹,train=True 表示下载训练集 train_set = torchvision.datasets.CIFAR10( root="./data", train=True, download=True, transform=None ) print("训练集样本数:", len(train_set))root="./data"指定数据落在当前目录的data子文件夹里,download=True表示本地没有时自动从官网下载,train=True拉取的是5万张训练图片。如果你只想验证流程,可以另跑一次train=False下载测试集。下载完成后,data目录下会出现cifar-10-batches-py文件夹,里面是data_batch_1到data_batch_5以及test_batch。手动检查时,只要这几个文件存在,基本就说明数据完整了。
这里有个容易忽略的点:代码里可能写的是相对路径"./data",所以你必须从项目根目录运行命令,否则程序会在错误的地方找不到数据。这也是我在开头强调先看README.md的原因——它通常会写明数据集该放哪个位置。
2.3 配置与加载:conf.json 与 datasets.py 的自适应逻辑
数据集就位后,配置文件就派上用场了。utils/conf.json是整套代码的中枢,它决定了服务端地址、客户端数量、全局轮次、本地训练轮次、批量大小和模型类型等关键参数。我一般会先打开这个文件,逐字段确认默认值是否合理,再启动训练。
{ "server_ip": "127.0.0.1", "server_port": 8080, "clients_num": 2, "global_rounds": 10, "local_epochs": 1, "batch_size": 32, "learning_rate": 0.001, "data_dir": "./data", "model": "cnn" }server_ip和server_port是服务端监听的地址,用127.0.0.1表示本机模拟联邦学习场景;clients_num是客户端数量,最少2个才能体现参数聚合;global_rounds是全局通信轮次,local_epochs是每个客户端本地训练的轮数;learning_rate控制本地优化器的学习率;data_dir指向2.2节准备好的数据集。如果配置里缺少data_dir这个键,datasets.py可能会默认拼一个./data路径,所以保持配置文件与代码约定一致很重要。
datasets.py负责把CIFAR-10按客户端编号切分。横向联邦的标准做法是每个客户端只拿到一部分样本,而不是所有数据。常见的切分逻辑是先把数据集随机打乱,再按clients_num分成若干份,每个客户端加载属于自己的那份。
def load_partition(dataset, client_id, clients_num): """把数据集按客户端数量均匀切分,返回指定客户端的数据子集""" indices = list(range(len(dataset))) # 每个客户端分到的样本数 = 总样本数 / 客户端数 per_client = len(dataset) // clients_num start = client_id * per_client end = start + per_client # 使用 Subset 按索引切分,避免重复拷贝图片数据 from torch.utils.data import Subset return Subset(dataset, indices[start:end])per_client表示每个客户端应该拿到的样本数量,client_id乘上per_client得到起始索引,start:end切片得到该客户端的索引范围。Subset是PyTorch提供的子集工具,它只保存原数据集的引用和索引列表,不会在内存里复制一份图片,比手动切片更省内存。切分时要注意clients_num必须能被样本总数整除,否则最后一个客户端拿到的样本数会和其他客户端不一致,影响聚合的公平性。
3. 核心代码结构与联邦通信机制:从入口到聚合的完整链路
3.1 参数入口:argparse_demo.py 与 main.py 如何把配置传给后续模块
argparse_demo.py是一个参数解析示例文件,作者用它演示Python标准库argparse的用法。实际启动代码是main.py,它读取-c参数指向的conf.json,把配置加载成字典后分发给服务端和客户端。
import json import argparse def load_config(config_path): """加载JSON配置文件并返回参数字典""" with open(config_path, "r", encoding="utf-8") as f: config = json.load(f) return config if __name__ == "__main__": parser = argparse.ArgumentParser(description="横向联邦图像分类启动入口") parser.add_argument("-c", "--config", type=str, required=True, help="配置文件路径,例如 ./utils/conf.json") args = parser.parse_args() config = load_config(args.config) print("配置文件加载完成,参数如下:") print(config)argparse.ArgumentParser负责解析命令行参数,-c和--config是同一个参数的两个写法,type=str表示接收字符串,required=True表示不传就报错。args.config拿到的是配置文件的路径,再交给load_config用json.load读成Python字典。这样设计的好处是:你想改任何超参数都不需要动代码,只改conf.json就行;换一组实验参数也只需复制一份配置,比在代码里写死方便得多。
3.2 模型与数据侧:models.py 里的CNN结构与datasets.py的数据划分
models.py定义了图像分类模型。这个项目是从零实现,所以作者没有直接用PyTorch官方预训练模型,而是写了一个小型的卷积神经网络。它的结构通常是“卷积层 + 激活函数 + 池化层 + 全连接层”的经典组合,输入是32×32的彩色图片,输出是10个类别的概率分布。
import torch.nn as nn class SimpleCNN(nn.Module): """用于CIFAR-10分类的简易卷积网络""" def __init__(self, num_classes=10): super(SimpleCNN, self).__init__() # 第一个卷积块:3通道输入,32通道输出 self.conv1 = nn.Sequential( nn.Conv2d(3, 32, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2) ) # 第二个卷积块:32通道转64通道 self.conv2 = nn.Sequential( nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2) ) # 全连接分类层:8*8*64 维输入,10维输出 self.fc = nn.Linear(8 * 8 * 64, num_classes) def forward(self, x): x = self.conv1(x) x = self.conv2(x) x = x.view(x.size(0), -1) return self.fc(x)nn.Conv2d(3, 32, kernel_size=3, padding=1)第一个参数3是输入通道数(RGB三通道),32是输出通道数。padding=1保证卷积后特征图尺寸不变,MaxPool2d(2, 2)把尺寸减半,经过两次池化后32×32的输入变成8×8。x.view(x.size(0), -1)是展平操作,把特征图拉成一维向量再送入全连接层。num_classes=10对应CIFAR-10的10个类别,你也可以改成其他分类任务。
datasets.py里的数据划分逻辑承担的是“横向切分”职责。横向联邦中,每个客户端的数据特征空间一致,但样本不同;load_partition把完整数据集切分给不同客户端,模拟了真实场景中不同机构各自持有部分数据的情况。配合2.3节里的Subset写法,程序就能做到不复制图片数据、只切索引。
3.3 服务端与客户端:server.py的聚合逻辑与client.py的本地更新
服务端和客户端是整个联邦学习的核心。server.py负责启动监听、接收各客户端的模型参数、按联邦平均算法(FedAvg)聚合,并把聚合后的全局模型下发;client.py负责加载本地数据、接收初始化的全局模型、在本地数据上训练若干轮,再把更新后的参数回传。
# server.py 核心聚合逻辑 def fed_avg(global_w, client_weights, clients_num): """对多个客户端的模型参数做加权平均""" averaged = {} for key in global_w.keys(): # 累加所有客户端的同名字典参数,再除以客户端数量 averaged[key] = sum(w[key] for w in client_weights) / clients_num return averagedglobal_w.keys()拿到的是模型所有参数层的名字,比如conv1.0.weight、fc.bias等。这段代码遍历每一层参数,把各个客户端回传的同层参数相加再除以客户端数量,得到新的全局参数。这里默认每个客户端权重相等,所以clients_num是分母。真实场景中如果各客户端样本量差异大,应该按样本比例加权平均,但作为入门代码,等权聚合更能突出核心思想。
客户端本地更新的逻辑在client.py里,本质就是标准的PyTorch训练过程:加载本地数据,把服务端下发的参数装进模型,用交叉熵损失和随机梯度下降优化器迭代若干轮。
# client.py 本地训练简化逻辑 import torch import torch.nn as nn def local_train(model, train_loader, lr, local_epochs): """在客户端本地数据上训练若干轮,返回更新后的模型参数""" criterion = nn.CrossEntropyLoss() optimizer = torch.optim.SGD(model.parameters(), lr=lr) model.train() for epoch in range(local_epochs): for images, labels in train_loader: optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() return model.state_dict()local_epochs是本地训练轮数,这个值设太大会导致客户端模型过拟合局部数据,设太小则参数更新不明显;lr是学习率,联邦场景中通常比单机训练要小,因为多次平均会放大学习率的影响。model.state_dict()返回的是模型全部参数组成的字典,这个字典就是客户端要回传给服务端的东西。整个过程可以看出,原始图片数据始终没有离开客户端,服务端拿到的只有梯度更新后的参数,这正是联邦学习保护数据隐私的关键。
3.4 main.py如何串起服务端与客户端
main.py是总调度器。它的职责是解析conf.json,然后启动服务端线程,再创建若干个客户端进程。每个客户端拿到自己的数据切片、初始化模型参数,开始与服务器通信。
# main.py 调度逻辑(示例结构) from server import FederatedServer from client import FederatedClient config = load_config("./utils/conf.json") server = FederatedServer(config) clients = [FederatedClient(config, client_id=i) for i in range(config["clients_num"])] for round_idx in range(config["global_rounds"]): # 每轮先让各客户端本地训练,再汇总聚合 weights = [client.local_train_and_send() for client in clients] new_weights = server.fed_avg(weights) server.broadcast(new_weights)从这段调度可以看到,global_rounds是全局通信的轮次,每轮由“各客户端本地训练 → 服务端聚合 → 广播回新参数”三步组成。理解这条循环,就理解了整个横向联邦学习的骨架:本地计算和中心聚合交替进行。调试时如果发现模型不收敛,先检查global_rounds是否足够、clients_num数据切分是否正确,都比改模型结构更有效。
4. 从零跑通训练:命令行启动、参数调优与结果观测
4.1 一行命令启动:python main.py -c ./utils/conf.json
环境就绪、数据落盘、代码结构也捋顺了,就可以真正跑训练。作者在摘要中给出的启动命令是:
python main.py -c ./utils/conf.json这条命令的意思是:用main.py启动整个联邦训练流程,用-c指定配置文件路径为./utils/conf.json。命令需要在项目根目录下执行,因为代码里的相对路径,比如./data、./utils/conf.json,都是相对于当前工作目录来解析的。如果你在别的目录执行命令,程序会报FileNotFoundError,找不到配置文件。这是新手最常踩的坑,不是代码问题,是工作目录问题。
启动后,正常情况会先看到配置文件加载的打印信息,然后是服务端启动、监听地址和端口,接着是各客户端连接上来的日志。训练过程中会定期输出当前的全局轮次、损失值和准确率。如果所有日志停在“服务端等待客户端连接”这一步,多半是客户端数量没凑齐,或者客户端进程启动失败,需要在另一个终端手动启动对应数量的客户端脚本。
4.2 训练参数的开放与调优:conf.json里每个字段怎么改
conf.json里的参数直接决定了训练行为和耗时。下表按“作用”和“调整建议”整理了一份,方便你在改配置时对照参考。
| 字段 | 作用 | 调整建议 |
|---|---|---|
server_ip/server_port | 服务端监听地址与端口 | 本机模拟用127.0.0.1,换端口需确保未被占用 |
clients_num | 模拟参与训练的客户端数量 | 建议2~4,数据集小或机器配置低时别设太大 |
global_rounds | 全局参数聚合轮次数 | CIFAR-10上10~30轮能看到明显趋势 |
local_epochs | 每个客户端本地训练轮数 | 1~3即可,过大易导致各客户端模型分叉 |
batch_size | 本地训练批量大小 | CPU环境设32;GPU可设64或128 |
learning_rate | 本地SGD优化器学习率 | 联邦场景建议0.001~0.01 |
model | 使用的模型名称标识 | 与models.py里注册的模型名对应上 |
clients_num直接影响每个客户端分到的数据量。如果设4,那么每个客户端只有1.25万张训练图,本地训练一轮的耗时和效果都会变化;global_rounds设太大训练时间会成倍增加,设太小则聚合效果不明显。我的建议是先用global_rounds=3跑通流程,确认无报错后再加大轮次。
还有几个参数在配置文件里可能不出现,但代码里会写死默认值,比如随机种子seed。如果需要实验可复现,最好在配置里补上并固定:
"seed": 2024固定随机种子后,数据打乱顺序、模型初始化的随机权重都会保持一致,多次运行结果可对比。不设种子的话,两次运行的初始模型不同,聚合结果的比较就没有意义。
4.3 结果观测与模型复用:从训练日志到准确率判断
训练完成后,项目根目录或figures文件夹里会出现一些图表,fig31.png、fig2.png就是作者在实验后生成的准确率和损失曲线图。这些图通常由matplotlib绘制,显示了全局模型在测试集上的表现随通信轮次变化的情况。
判断训练是否正常,核心看两个指标:一是损失是否随global_rounds下降,二是测试集准确率是否上升。横向联邦与集中式训练的区别在于,由于客户端各自只看到部分数据,且不断经历“本地训练 → 参数平均”的过程,损失曲线会有明显的抖动,不像单机训练那样平滑。只要整体趋势向下,就说明聚合逻辑是正确的。
过程中要留意客户端配置与数据量的匹配。如果每个客户端数据太少而local_epochs很大,模型会对局部数据过拟合,聚合后的全局模型表现反而变差。遇到这种情况,优先调小local_epochs或调大clients_num对应的每客户端样本量,而不是盲目增加全局轮次。代码包里的figures目录可以作为你实验结果的参照物,跑出来的曲线形状接近,基本就说明操作没问题。
5. 避坑与常见问题排查:四个踩过才会懂的坑
5.1 现象:运行报错“FileNotFoundError”,提示找不到CIFAR-10数据或配置文件
新手最容易遇到的就是路径问题。明明下载了数据集,程序还是报找不到cifar-10-batches-py,或者提示./utils/conf.json不存在。
原因:当前工作目录不在项目根目录,或者数据集没有解压到data文件夹的指定层级。torchvision下载的数据要放在data/cifar-10-batches-py/目录下,如果把压缩包直接扔在data里就认为完事,程序当然找不到。
解决:先cd到项目根目录再执行命令,用pwd确认当前路径;再用ls data检查是否出现cifar-10-batches-py文件夹。代码里的data_dir是相对路径时,务必保持控制台目录与配置一致。我习惯在main.py开头打印os.getcwd(),一跑就知道有没有跑错目录,这是最直接的排查方式。
5.2 现象:服务端一直卡在“等待客户端连接”,或启动时报“Address already in use”
服务端启动报端口被占用,或者客户端怎么也连不上。
原因:上一次异常退出后,服务端进程没有完全释放端口;或者server_port和本机其他程序冲突。Windows下进程未正常结束、僵尸Python进程残留,都会导致8080端口被占用。
解决:先找占用端口的进程,再决定是换端口还是杀掉旧进程。
# Linux / macOS 下查看 8080 端口占用情况 lsof -i :8080 # 找到 PID 后结束对应进程 kill -9 <PID>Windows下可以用netstat -ano | findstr 8080查到PID,再用taskkill /PID <PID> /F强制结束。嫌麻烦就换一个端口,比如9090,然后同步改掉conf.json里的server_port。模拟联邦学习时,所有客户端连接同一个服务端地址和端口,改配置后要确保所有地方都用了新端口。
5.3 现象:训练过程报错“CUDA out of memory”或“Expected all tensors to be on the same device”
训练跑到一半爆显存,或者报设备不一致的错。
原因:默认模型参数和数据被放在不同设备上。CIFAR-10图片先被加载到CPU,模型参数却跑到了GPU,或者反过来。如果你的机器显存只有4GB,batch size设太大也会爆显存。
解决:在client.py本地训练前,把模型和数据都显式搬到同一设备。常见做法是加一个设备判断:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = model.to(device) # 训练循环里把数据也搬到同一设备 images, labels = images.to(device), labels.to(device)torch.cuda.is_available()检查GPU是否可用,没有GPU就自动退回CPU,代码两端都兼容。显存溢出时,优先把batch_size从64改到32甚至16,这个参数对显存占用影响最直接。
5.4 现象:聚合后模型不收敛,准确率比单一客户端本地训练还低
训练能跑完,但全局模型的准确率一直上不去,甚至不如随便选一个客户端单独训练的结果。
原因:非独立同分布数据划分不当,或者客户端数量与local_epochs搭配失衡。如果数据切分时没有全局打乱,每个客户端拿到的可能是某一类别的集中样本;本地训练轮次过大又会让客户端模型偏离全局方向。
解决:检查datasets.py里是否先对整个数据集做了random.shuffle或用了torch.utils.data.random_split再切分。确保每个客户端都包含所有类别的样本,这是横向联邦的基本前提。再把local_epochs降回1,learning_rate调小到0.001,让每个客户端每次只走一小步,聚合才能平滑稳定。我曾在这上面耗过整整一天——模型、聚合代码都没问题,就因为是按顺序切数据,前1/2全是同一批类别,聚合效果自然崩。
6. 进阶调试技巧:用迷你数据集验证聚合逻辑并深挖代码
6.1 构造迷你CIFAR子集,快速验证一轮联邦训练
完整CIFAR-10有5万张训练图,调一次参要等很久。我习惯先构造一个迷你子集,比如每个类别取200张,把训练集压缩到2000张,专门用来验证代码链路是否畅通。
from torch.utils.data import Subset import torchvision # 取前2000个样本作为调试子集 full_set = torchvision.datasets.CIFAR10( root="./data", train=True, download=False, transform=None ) mini_indices = list(range(2000)) mini_set = Subset(full_set, mini_indices)download=False表示不再下载数据,Subset按索引切片出前2000条。用这个迷你集替换datasets.py里的数据源,训练一轮只要几秒,可以快速验证服务端和客户端的参数是否真的在聚合。确认逻辑没问题后,再换回完整数据集跑正式实验。
6.2 在server.py聚合处加临时调试代码,直观观察参数变化
只看准确率曲线,你很难确认“聚合”到底发生了什么。我会在fed_avg聚合函数里临时加一段打印,观察每一层参数在聚合前后的变化幅度:
# 调试用:观察某一层参数聚合前后的差异 layer_name = "fc.weight" import numpy as np before = global_w[layer_name].clone() # 执行聚合... after = averaged[layer_name] diff = (before - after).abs().mean().item() print(f"{layer_name} 平均变化幅度: {diff:.6f}")before - after计算聚合前后参数差,abs().mean()得到平均绝对变化量。每一轮打印这个数值,如果变化幅度逐渐减小,说明模型正在收敛;如果数值剧烈跳动且不下降,说明学习率太大或数据划分有问题。这段代码在正式实验里可以删掉,但调试时它比任何日志都好用。
6.3 扩展模型:把models.py的CNN换成ResNet
理解现有CNN后,你可能会想把模型换成更深的ResNet。models.py里只需要把现有的SimpleCNN替换成PyTorch自带的resnet18。
import torchvision.models as models def create_model(model_name, num_classes=10): """根据模型名称创建分类模型""" if model_name == "resnet18": net = models.resnet18(weights=None, num_classes=num_classes) else: net = SimpleCNN(num_classes=num_classes) return netweights=None表示不加载预训练权重,因为联邦学习场景必须保证所有客户端从同一个初始模型开始训练。num_classes=10把ResNet最后的全连接层改成10类输出。注意datasets.py里原来的数据变换可能只适用于32×32的小图,ResNet经典输入是224×224,需要在加载CIFAR-10时补一个Resize操作,否则模型的前向传播会报尺寸不匹配。
从那以后,我每次跑联邦学习实验都强制自己先走一遍迷你数据集加单轮聚合调试,确认参数真的动了、数据切分没漏样本,再上完整训练。这套顺序成型后,我几乎没有再被“聚合不收敛”“端口被占用”这类问题卡过半天以上。代码这种东西,自己动手改一次,比看十遍注释都管用。希望这份横向联邦学习代码的拆解能帮到你,也欢迎在复现过程中翻看源码里的大量注释——那些注释才是作者留下的真正财富。
本文还有配套的精品资源,点击获取