☰
横向联邦图像分类从零实现:FedAvg与PyTorch实战指南
2026/10/7 6:21:02 网站建设 项目流程

简介:代码项目对应《联邦学习实战》第3章横向联邦图像分类的完整配套实现,面向需要入门联邦学习与图像分类的Python开发者、高校学生及竞赛人员,可在PyTorch环境中结合CIFAR-10数据集直接运行。压缩包内有22个文件,约156KB,主要包含Python源文件、编译缓存pyc、项目配置xml/json、数据集目录占位、示意图和说明文档等;代码划分为主程序、服务端、客户端、数据集加载、模型定义等模块,每个模块均附有大量注释,注释覆盖参数配置、通信交互与模型更新的关键逻辑,能帮助读者快速理解横向联邦的基本流程。目前已有126人学习下载,尤其适合计算机、人工智能、自动化等相关专业用于毕设、课设、课程作业或自学进阶。资源已通过实际运行验证,回应了环境安装、数据集放置和启动方式等常见问题,下载后可按README说明快速搭建运行,也可在现有代码基础上做二次改进,节省不少前期排错时间。

1. 横向联邦图像分类从零写起:为什么我建议先别碰联邦框架

想学横向联邦图像分类的人,第一反应往往是去装一个联邦学习框架,然后跑通官方示例。实际做下来你会发现,你连“客户端数据是怎么切的”“服务端到底聚合了什么”都没搞清楚,框架就把活干完了。横向联邦图像分类的本质,是让多个客户端各自持有本地图片和标签,协同训练同一个分类模型,服务端只聚合模型权重、不碰原始图像;而“基于python从零实现”,意味着你只用最朴素的 PyTorch 和 numpy 就能把这条链路拼出来。这篇笔记适合已经能跑通普通图像分类、想进入联邦学习但不想被框架黑匣子劝退的人。我会按数据切分、客户端训练、服务端聚合、踩坑、进阶实验的顺序,把一份带大量注释的学习代码拆给你看。

2. FedAvg 与图像分类模型选型:聚合公式、小 CNN 结构与通信成本

2.1 横向联邦为什么是图像分类的最佳入门场景

横向联邦学习里,数据按“样本维度”切分:每个客户端拥有的是不同的图片样本,但类别空间完全一致,比如大家都在分“猫、狗、飞机、汽车”。图像分类恰好是这个设定下最顺手的任务,因为它的损失函数就是交叉熵,计算图清晰,本地训练和单机训练没有本质差别。你不需要处理跨域特征对齐、不需要设计标签体系映射,只需要回答三个问题:数据怎么分、模型怎么传、权重怎么合。

这也是为什么医院影像分诊、手机端相册分类这类业务愿意用横向联邦:每个机构的数据格式统一,只是样本不互通。与之相对,森林图像分类那种专业场景虽然也是图像任务,但客户端之间的拍摄条件、物种分布差异性太大,入门阶段很难判断“准确率上不去”到底是联邦算法问题还是数据问题。所以我强烈建议入门数据用 CIFAR-10 或 MNIST,而不是一上来就挑战大而偏的专业图像集。

2.2 FedAvg 聚合公式与一个最小的 numpy 版本

横向联邦里最经典的算法是 FedAvg。假设有 N 个客户端,第 k 个客户端持有 n_k 张图片,总样本数 n = Σ n_k。每一轮通信的过程可以拆成四步:服务端把当前全局模型 w_t 分发给选中的客户端;每个客户端用本地数据跑若干个 epoch 的 SGD,得到本地模型 w_{t+1}^k;客户端把权重传回服务端;服务端按样本量加权平均,生成新的全局模型 w_{t+1}。聚合公式是:

w_{t+1} = Σ ( n_k / Σ n_j ) · w_{t+1}^k

这里的关键是“按样本数加权”而不是简单平均。样本多的客户端见过更多图像,它的本地模型在经验上更可信,权重应该更大。这个直觉用一个最小 numpy 函数就可以验证:

import numpy as np def fedavg_aggregate(client_weights, client_sizes): # client_weights: 每个参与客户端的模型参数展平后的数组 # client_sizes: 每个参与客户端的本地样本数 total = sum(client_sizes) avg_weight = np.zeros_like(client_weights[0]) for w, size in zip(client_weights, client_sizes): avg_weight += w * (size / total) return avg_weight

这段代码的逻辑很简单:先算总样本数,然后每个客户端的权重向量乘上它的样本占比,累加就是新的全局参数。参数上要注意两点:client_weights里的每个数组必须是从同一个全局模型出发训练得到的,否则逐元素相加没有数学意义;client_sizes建议用“参与本轮客户端”的样本数重新计算占比,而不是全局总样本数。很多初学者在这里偷懒用简单平均,非 IID 数据下收敛速度会明显变慢。

2.3 图像分类模型选型:小 CNN 的三条取舍标准

图像分类模型怎么选,是这份学习代码里最容易被忽视的一步。我的建议是:不要用最新的图像分类模型,不要用预训练大模型,就手写一个两层卷积的小 CNN。原因有三个。第一是通信成本,联邦学习每一轮都要把完整的 state_dict 从服务端传到客户端再传回来,模型参数翻一倍,通信时间就翻一倍,大模型在入门阶段会让你把时间耗在等训练上。第二是客户端异构性,真实场景里客户端算力差距很大,小模型更容易在各端跑完本地训练。第三是可读性,你是在学联邦,不是在学模型结构,模型越简单,问题越容易定位到联邦机制本身。

import torch import torch.nn as nn class SmallCNN(nn.Module): def __init__(self, num_classes=10): super().__init__() # CIFAR-10 输入是 3x32x32,先用两组卷积提取特征 self.features = nn.Sequential( nn.Conv2d(3, 32, kernel_size=3, padding=1), # 输出 32x32x32 nn.ReLU(inplace=True), nn.MaxPool2d(2), # 输出 32x16x16 nn.Conv2d(32, 64, kernel_size=3, padding=1), # 输出 64x16x16 nn.ReLU(inplace=True), nn.MaxPool2d(2), # 输出 64x8x8 ) self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(64 * 8 * 8, 128), nn.ReLU(inplace=True), nn.Linear(128, num_classes), ) def forward(self, x): return self.classifier(self.features(x))

这个模型总共约 60 万参数,在 CPU 上也能跑得动。结构上我特意没有加 BN 层,原因在后面的避坑章节会展开。用padding=1是为了保持特征图尺寸,这样全连接层的输入尺寸好算。如果你想让训练更稳,可以把nn.ReLU换成nn.GELU,但入门阶段不需要。参数上最需要记住的是num_classes=10,如果你的数据集不是 CIFAR-10,记得把这里和下游的数据集类别数对齐。

3. 从零实现横向联邦图像分类代码:数据切分、客户端训练与服务端聚合

3.1 python 环境准备与目录结构

开始写代码之前,先把 python 环境配好。这份学习代码依赖torch、torchvision、numpy和matplotlib,python 版本 3.8 到 3.11 都可以。安装命令建议用 pip 一次性装完,如果你之前装过 torch 但不确定版本,先pip show torch看一眼,别混装 CPU 版和 GPU 版,这是 python 环境配置里最常见的翻车点。

pip install torch torchvision numpy matplotlib tqdm

目录结构我建议按职责拆成五个文件,而不是把全部逻辑堆在一个脚本里。这样你后续加差分隐私、换数据集、改聚合方式,都只需要动对应文件。文件划分如下:

文件职责
dataset.py下载并加载 CIFAR-10,实现非 IID 数据切分
model.py定义 SmallCNN 图像分类模型
client.py定义客户端类,实现本地训练
server.py定义服务端类,实现 FedAvg 聚合与测试集评估
main.py主训练循环,串起数据、客户端、服务端

从零实现的核心原则是“一个文件只做一件事”。很多初学者把客户端训练和服务端聚合写在同一个类的同一个方法里,改参数时牵一发动全身。我的习惯是先写main.py里的主循环伪代码,再回头补每个类的实现,这样整体流程能在头脑里先跑通。

3.2 非 IID 数据切分:用 Dirichlet 分布模拟真实客户端偏移

横向联邦的难点不在模型,而在数据。真实场景里每个客户端的图片分布几乎不可能一致:有的客户端只有猫狗,有的客户端只有汽车飞机。这种“非 IID”数据分布如果不用代码模拟,你写出来的联邦代码在测试里再漂亮,落地也会原形毕露。常见做法是用 Dirichlet 分布按类别生成每个客户端的样本比例,alpha参数控制偏移程度,alpha越小分布越偏。

import numpy as np def dirichlet_split(dataset, num_clients, alpha=0.5): # dataset: torchvision 的 CIFAR-10 数据集 # num_clients: 模拟的客户端数量 # alpha: Dirichlet 分布参数,越小数据越不均衡 targets = np.array(dataset.targets) num_classes = len(np.unique(targets)) # 为每个类别生成它在各客户端上的占比,shape: num_classes x num_clients ratios = np.random.dirichlet([alpha] * num_clients, size=num_classes) class_indices = [] for c in range(num_classes): cidx = np.where(targets == c)[0] np.random.shuffle(cidx) # 按占比计算当前类别每个客户端应分到的样本数 split_points = (np.cumsum(ratios[c]) * len(cidx)).astype(int)[:-1] class_indices.append(np.split(cidx, split_points)) # 按客户端维度聚合所有类别 client_indices = [ np.concatenate([class_indices[c][i] for c in range(num_classes)]) for i in range(num_clients) ] for ci in client_indices: np.random.shuffle(ci) return client_indices

逻辑上这段代码先按类别把全部样本分堆,再在每个类别内部按 Dirichlet 比例切给不同客户端,最后按客户端维度拼起来。这样做的好处是每个类别都参与了分配,不会出现某个类别整体丢失的情况。参数上alpha=0.5是一个适中的非 IID 程度,alpha=100时分布接近均分,alpha=0.1时很多客户端会只剩一到两个类别。

注意:np.split要求切分点严格递增且不能超出数组长度。如果你的客户端数量很多、某些类别样本很少,split_points里可能出现重复值甚至越界,建议在切分前把split_points去重,并对长度不足的类别做轮询补样。

3.3 客户端本地训练:从全局权重出发,而不是从随机权重出发

客户端类是整个联邦代码里最容易写错的地方。核心细节是:每一轮训练开始前,客户端必须无条件加载服务端下发的全局权重,而不是沿用上一轮自己的本地权重。这是联邦和分布式训练的本质区别——分布式训练各节点算同一个目标,联邦各客户端算的是有偏的本地目标,如果客户端自说自话延续本地状态,全局模型很快会被拽偏。

import torch import torch.nn as nn class Client: def __init__(self, cid, train_loader, model_fn, device="cpu"): self.cid = cid self.train_loader = train_loader self.model = model_fn().to(device) self.device = device def set_global_model(self, global_state_dict): # 必须调用 load_state_dict,覆盖本地旧参数 self.model.load_state_dict(global_state_dict) def local_update(self, epochs=5, lr=0.01, momentum=0.9): # 本地用普通 SGD 训练若干个 epoch optimizer = torch.optim.SGD(self.model.parameters(), lr=lr, momentum=momentum) criterion = nn.CrossEntropyLoss() self.model.train() total_loss = 0.0 num_batches = 0 for _ in range(epochs): for images, labels in self.train_loader: images, labels = images.to(self.device), labels.to(self.device) optimizer.zero_grad() outputs = self.model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() total_loss += loss.item() num_batches += 1 # 返回更新后的模型权重和本轮平均损失 return self.model.state_dict(), total_loss / max(num_batches, 1)

set_global_model是每轮训练前必须调用的方法,这一步漏了,你的联邦学习就退化成 N 个互不相干的单机训练。local_update里的epochs是本地训练轮数,不是全局通信轮数,初学者经常把这两个搞混。参数上我推荐的起点是epochs=5、lr=0.01、momentum=0.9,这个组合在 CIFAR-10 上能在 50 轮通信内看到明显收敛。如果你的客户端数量很多,epochs可以降到 1 或 2,避免单客户端过拟合本地数据。

3.4 服务端 FedAvg 聚合与主训练循环

服务端类负责两件事:按样本占比聚合客户端权重,以及在聚合后用服务端持有的测试集评估全局模型。这里有一个设计决定要提前想清楚:服务端要不要保留一份干净测试集。我的做法是保留,因为学习代码需要一个统一的评估口径来判断每轮通信是否进步,否则你只能看客户端上报的本地 loss,数据偏的情况下这个数字毫无参考价值。

import torch class FedServer: def __init__(self, global_model, test_loader, device="cpu"): self.global_model = global_model.to(device) self.test_loader = test_loader self.device = device def aggregate(self, client_updates, client_sizes): # client_updates: 每个参与客户端的 state_dict 列表 # client_sizes: 每个参与客户端的样本数列表 keys = client_updates[0].keys() total = sum(client_sizes) new_state = {} for k in keys: # 按样本占比加权求和,等价于 FedAvg 公式 new_state[k] = sum( update[k].float() * (size / total) for update, size in zip(client_updates, client_sizes) ) self.global_model.load_state_dict(new_state) def evaluate(self): self.global_model.eval() correct, total = 0, 0 with torch.no_grad(): for images, labels in self.test_loader: images, labels = images.to(self.device), labels.to(self.device) preds = self.global_model(images).argmax(dim=1) correct += (preds == labels).sum().item() total += labels.size(0) return correct / total

聚合时我按 state_dict 的每个 key 分别做加权求和,这是因为conv.weight、fc.bias这些张量形状不同,不能直接对整个 state_dict 做张量乘法。注意load_state_dict里的权重必须和原始模型结构严格对齐,这也是我在model.py里固定模型结构的原因。主训练循环放在main.py里,每轮随机抽取一部分客户端参与训练,抽取比例用client_ratio控制:

import numpy as np def run_federated(args, clients, server): for rnd in range(args.rounds): # 每轮随机抽取 client_ratio 比例的客户端参与 num_sampled = max(1, int(args.client_ratio * len(clients))) sampled = np.random.choice(clients, size=num_sampled, replace=False) updates, sizes = [], [] for client in sampled: # 关键:先下发全局模型,再做本地训练 client.set_global_model(server.global_model.state_dict()) state, loss = client.local_update(epochs=args.local_epochs, lr=args.lr) updates.append(state) sizes.append(len(client.train_loader.dataset)) server.aggregate(updates, sizes) acc = server.evaluate() print(f"round {rnd}: global acc = {acc:.4f}")

这个循环就是整个联邦学习的骨架。client_ratio=0.5表示每轮只有一半客户端参与,这是为了让代码更接近真实联邦场景——真实系统里客户端可能随时掉线。如果你的目的是验证算法收敛性,可以把client_ratio设成 1.0,让所有客户端每轮都参与,收敛会更稳定。

4. 横向联邦图像分类的 5 个踩坑现场:不收敛、BN 波动与日志泄露

4.1 全局准确率常年不动:优化器参数与非 IID 震荡

现象:跑了 30 轮通信,全局测试准确率一直停在 20% 到 30% 之间,偶尔还往下掉。原因有两层:一是lr=0.01对 CIFAR-10 这种任务本身偏高,二是非 IID 数据下各客户端的局部梯度方向差异很大,全局模型被拉来拉去,形成震荡。解决方法是把学习率降到0.001,本地训练轮数从 5 降到 1 或 2,并把客户端参与比例从 0.5 提到 0.8。这三个改动本质上是让全局模型每轮只走一小步,不被任何单一客户端带偏。

4.2 少数类别被客户端吞掉:Dirichlet 空类问题

现象:某个客户端本地准确率很高,但服务端全局模型对某个类别的预测永远错误。原因是我在 3.2 节提到的切分问题——alpha很小时,某个客户端可能完全分不到某个类别的样本,本地模型对那个类别的决策边界完全失效。解决方法是在切分后做一个空类回填:统计每个客户端拥有的类别集合,对缺失的类别,从全局该类别的样本池里补抽几条进去,保证每个客户端至少有 2 个类别的样本。这一步不是可选项,alpha=0.1下几乎必然触发。

4.3 BN 统计量在联邦中的不稳定:换归一化层或服务端重算

现象:模型结构里加了nn.BatchNorm2d之后,全局测试准确率出现明显抖动,而且客户端本地训练时 loss 正常、评估时却很差。原因是 BN 层的running_mean和running_var是在本地数据上累计的,非 IID 数据下每个客户端算出的统计量差异巨大,聚合时这些统计量被平均成了一个不伦不类的中间值。解决方法是二选一:把 BN 换成nn.GroupNorm,或者训练结束后在服务端用少量干净数据重算 BN 统计量。我给的SmallCNN里直接不用归一化层,就是为了绕开这个问题。

4.4 全局模型悄悄退化:检查初始化链路

现象:训练了 50 轮,某一天你发现第 20 轮的全局权重文件比第 50 轮的准确率还高,整个训练过程像是在原地打转。原因是Client.set_global_model没有被调用,或者调用时传错了对象——客户端一直在从随机初始化权重开始训练,服务端聚合的其实是 N 个互不相关的随机模型。排查方法很简单:训练前后打印sum(p.sum() for p in model.parameters())的数值,如果客户端在加载全局权重前后这个值没有变化,说明加载链路断了。

4.5 客户端日志泄漏:区分学习代码与生产协议

现象:为了调试方便,你把每个客户端的本地模型权重分别存成了client_7_round_12.pt,然后某一天意识到,这个文件落到别人手里,等于把客户端数据的信息通过模型权重泄露出去了。原因是在学习代码里养成了“直接保存每个客户端产物”的习惯。解决方法是在学习阶段养成两个好习惯:日志里只记录聚合后的全局模型信息,不记录单个客户端的梯度或权重;模拟隐私保护时,至少要在聚合前给权重加噪声,这部分在下一章展开。这个坑不会影响你的代码运行,但会影响你想不想把自己的代码用在真实数据上。

5. 让学习代码变成可信实验:非 IID 程度、参与比例与差分隐私模拟

5.1 用 alpha 把非 IID 程度变成可调旋钮

你在论文里会看到“在非 IID 设置下”这种说法,但很少有人告诉你非 IID 到底怎么量化。用 Dirichlet 分布的alpha参数就是最直观的旋钮。alpha越小,每个客户端的类别分布越极端;alpha越大,越接近均匀分布。我建议你在写完数据切分后,先做一张客户端类别分布图,确认你的“非 IID”符合直觉再开始训练。

import matplotlib.pyplot as plt import numpy as np def plot_client_distribution(client_indices, dataset, num_clients, save_path="client_dist.png"): # client_indices: dirichlet_split 返回的每个客户端的样本索引列表 targets = np.array(dataset.targets) num_classes = len(np.unique(targets)) matrix = np.zeros((num_clients, num_classes), dtype=int) for i, indices in enumerate(client_indices): for c in range(num_classes): matrix[i, c] = np.sum(targets[indices] == c) # 画堆叠条形图,每个柱子代表一个客户端 fig, ax = plt.subplots(figsize=(10, 4)) bottom = np.zeros(num_clients) for c in range(num_classes): ax.bar(range(num_clients), matrix[:, c], bottom=bottom, label=f"class {c}") bottom += matrix[:, c] ax.set_xlabel("client id") ax.set_ylabel("sample count") ax.legend() fig.savefig(save_path)

参数上,alpha=0.1会让多数客户端只剩一个主导类别,alpha=1.0是中等偏移,alpha=10以上基本接近 IID。我的血泪经验是:永远不要只跑一个alpha就下结论,至少跑0.1, 0.5, 1.0, 100四档,你才能判断你的联邦算法对数据偏移到底有多敏感。

5.2 参与比例 C 与通信轮次的实验矩阵

很多学习代码默认每轮所有客户端都参与,这在实验里是可行的,但会掩盖联邦系统的一个核心矛盾:客户端参与比例越低,每轮通信成本越低,但全局模型的收敛越不稳定。建议你跑一个 3x3 的小实验矩阵,把参与比例和通信轮次对应起来:

参与比例 C通信轮次预期收敛表现
0.350震荡明显,准确率波动大
0.5100基本收敛,非 IID 下有 2-3 个点波动
1.0100收敛最稳,但每轮耗时最长

跑实验时把随机种子固定住,让同一个(C, rounds)组合可以复现。如果你发现C=0.3的准确率比C=1.0只低 1 到 2 个点,说明你的数据切分偏 IID;如果差 5 个点以上,说明非 IID 程度已经影响到了全局模型的稳定性。这个对比本身就是一份很好的实验结论。

5.3 用裁剪与高斯噪声模拟差分隐私扰动

横向联邦的论文里经常出现“安全聚合”“差分隐私”这些词,学习代码不需要实现完整协议,但至少要模拟隐私保护对模型精度的影响。最简做法是在服务端聚合前,对每个客户端的权重做范数裁剪,然后加高斯噪声。

import torch def clip_and_noise_aggregate(client_updates, client_sizes, clip_norm=1.0, sigma=0.01): # client_updates: state_dict 列表,先裁剪再加噪再做加权平均 keys = client_updates[0].keys() total = sum(client_sizes) noised_state = {} for k in keys: # 按客户端逐参数裁剪到 clip_norm 范数以内 clipped = [] for update in client_updates: param = update[k].float() norm = param.norm() if norm > clip_norm: param = param * (clip_norm / norm) clipped.append(param) # 加权平均后叠加高斯噪声 avg = sum(p * (s / total) for p, s in zip(clipped, client_sizes)) noise = torch.randn_like(avg) * sigma noised_state[k] = avg + noise return noised_state

clip_norm控制单客户端权重的影响上限,sigma控制噪声强度。sigma越大,隐私保护越强,但全局模型准确率掉得越多。你可以把sigma从0.001调到0.05,画一条“隐私-精度”的下降曲线,这是联邦学习里最值得展示的实验结果之一。注意这里的实现只是模拟,真实差分隐私还需要按敏感度计算噪声尺度,但作为学习代码,理解“扰动发生在聚合前”这个时序就够了。

6. 用固定种子与指标落盘,把学习代码变成可复现实验

6.1 固定随机种子与实验记录清单

做到这一步,你的代码已经能跑通横向联邦图像分类的完整流程,但还有一个会毁掉所有实验的隐患:随机性。PyTorch、numpy、Python 自带 random 三套随机源,任何一套不固定,同一个参数跑两次结果都不一样。我见过最夸张的一次,同一个代码跑两遍,准确率差了 5 个点,原因就是 DataLoader 的 shuffle 和 Dirichlet 切分都没固定种子。修复方法很直接:

import random import numpy as np import torch def seed_everything(seed=42): # 固定 python、numpy、torch 三套随机源,保证实验可复现 random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False

cudnn.deterministic=True会牺牲一点性能,但换来的是卷积计算的确定性。benchmark=False禁止 cuDNN 在运行时自动选择算法,否则算法选择的随机性也会影响结果。之后把你的实验记录落盘成 CSV,每行一条:

import csv def log_round(row, path="fl_results.csv"): # row: {"round": 10, "acc": 0.7234, "lr": 0.001, "alpha": 0.5, "client_ratio": 0.5} with open(path, "a", newline="") as f: writer = csv.DictWriter(f, fieldnames=list(row.keys())) if f.tell() == 0: writer.writeheader() writer.writerow(row)

记录字段至少包括:当前轮次、全局测试准确率、学习率、本地 epoch 数、客户端参与比例、Dirichlet 的 alpha、随机种子、数据切分版本号。我自己的习惯是把这些信息直接拼进文件名,比如cifar10_a0.5_ratio0.5_lr0.001_r42_v3.csv,这样就算日志文件堆满一个文件夹,也不会搞混哪份实验对应哪组参数。这是我做联邦学习实验吃过亏之后养成的习惯——不固定种子、不落盘指标,你的代码跑出来的任何“结论”都只是巧合。希望这份笔记能帮你把横向联邦图像分类的学习代码跑通,并跑出能说服自己的结果。

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

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

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

立即咨询