横向联邦图像分类从零实现:FedAvg与PyTorch实战指南 简介代码项目对应《联邦学习实战》第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_{t1}^k客户端把权重传回服务端服务端按样本量加权平均生成新的全局模型 w_{t1}。聚合公式是w_{t1} Σ ( n_k / Σ n_j ) · w_{t1}^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_classes10): super().__init__() # CIFAR-10 输入是 3x32x32先用两组卷积提取特征 self.features nn.Sequential( nn.Conv2d(3, 32, kernel_size3, padding1), # 输出 32x32x32 nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # 输出 32x16x16 nn.Conv2d(32, 64, kernel_size3, padding1), # 输出 64x16x16 nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # 输出 64x8x8 ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 8 * 8, 128), nn.ReLU(inplaceTrue), nn.Linear(128, num_classes), ) def forward(self, x): return self.classifier(self.features(x))这个模型总共约 60 万参数在 CPU 上也能跑得动。结构上我特意没有加 BN 层原因在后面的避坑章节会展开。用padding1是为了保持特征图尺寸这样全连接层的输入尺寸好算。如果你想让训练更稳可以把nn.ReLU换成nn.GELU但入门阶段不需要。参数上最需要记住的是num_classes10如果你的数据集不是 CIFAR-10记得把这里和下游的数据集类别数对齐。3. 从零实现横向联邦图像分类代码数据切分、客户端训练与服务端聚合3.1 python 环境准备与目录结构开始写代码之前先把 python 环境配好。这份学习代码依赖torch、torchvision、numpy和matplotlibpython 版本 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, alpha0.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, sizenum_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 比例切给不同客户端最后按客户端维度拼起来。这样做的好处是每个类别都参与了分配不会出现某个类别整体丢失的情况。参数上alpha0.5是一个适中的非 IID 程度alpha100时分布接近均分alpha0.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, devicecpu): 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, epochs5, lr0.01, momentum0.9): # 本地用普通 SGD 训练若干个 epoch optimizer torch.optim.SGD(self.model.parameters(), lrlr, momentummomentum) 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是本地训练轮数不是全局通信轮数初学者经常把这两个搞混。参数上我推荐的起点是epochs5、lr0.01、momentum0.9这个组合在 CIFAR-10 上能在 50 轮通信内看到明显收敛。如果你的客户端数量很多epochs可以降到 1 或 2避免单客户端过拟合本地数据。3.4 服务端 FedAvg 聚合与主训练循环服务端类负责两件事按样本占比聚合客户端权重以及在聚合后用服务端持有的测试集评估全局模型。这里有一个设计决定要提前想清楚服务端要不要保留一份干净测试集。我的做法是保留因为学习代码需要一个统一的评估口径来判断每轮通信是否进步否则你只能看客户端上报的本地 loss数据偏的情况下这个数字毫无参考价值。import torch class FedServer: def __init__(self, global_model, test_loader, devicecpu): 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(dim1) 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, sizenum_sampled, replaceFalse) updates, sizes [], [] for client in sampled: # 关键先下发全局模型再做本地训练 client.set_global_model(server.global_model.state_dict()) state, loss client.local_update(epochsargs.local_epochs, lrargs.lr) updates.append(state) sizes.append(len(client.train_loader.dataset)) server.aggregate(updates, sizes) acc server.evaluate() print(fround {rnd}: global acc {acc:.4f})这个循环就是整个联邦学习的骨架。client_ratio0.5表示每轮只有一半客户端参与这是为了让代码更接近真实联邦场景——真实系统里客户端可能随时掉线。如果你的目的是验证算法收敛性可以把client_ratio设成 1.0让所有客户端每轮都参与收敛会更稳定。4. 横向联邦图像分类的 5 个踩坑现场不收敛、BN 波动与日志泄露4.1 全局准确率常年不动优化器参数与非 IID 震荡现象跑了 30 轮通信全局测试准确率一直停在 20% 到 30% 之间偶尔还往下掉。原因有两层一是lr0.01对 CIFAR-10 这种任务本身偏高二是非 IID 数据下各客户端的局部梯度方向差异很大全局模型被拉来拉去形成震荡。解决方法是把学习率降到0.001本地训练轮数从 5 降到 1 或 2并把客户端参与比例从 0.5 提到 0.8。这三个改动本质上是让全局模型每轮只走一小步不被任何单一客户端带偏。4.2 少数类别被客户端吞掉Dirichlet 空类问题现象某个客户端本地准确率很高但服务端全局模型对某个类别的预测永远错误。原因是我在 3.2 节提到的切分问题——alpha很小时某个客户端可能完全分不到某个类别的样本本地模型对那个类别的决策边界完全失效。解决方法是在切分后做一个空类回填统计每个客户端拥有的类别集合对缺失的类别从全局该类别的样本池里补抽几条进去保证每个客户端至少有 2 个类别的样本。这一步不是可选项alpha0.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_pathclient_dist.png): # client_indices: dirichlet_split 返回的每个客户端的样本索引列表 targets np.array(dataset.targets) num_classes len(np.unique(targets)) matrix np.zeros((num_clients, num_classes), dtypeint) 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], bottombottom, labelfclass {c}) bottom matrix[:, c] ax.set_xlabel(client id) ax.set_ylabel(sample count) ax.legend() fig.savefig(save_path)参数上alpha0.1会让多数客户端只剩一个主导类别alpha1.0是中等偏移alpha10以上基本接近 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)组合可以复现。如果你发现C0.3的准确率比C1.0只低 1 到 2 个点说明你的数据切分偏 IID如果差 5 个点以上说明非 IID 程度已经影响到了全局模型的稳定性。这个对比本身就是一份很好的实验结论。5.3 用裁剪与高斯噪声模拟差分隐私扰动横向联邦的论文里经常出现“安全聚合”“差分隐私”这些词学习代码不需要实现完整协议但至少要模拟隐私保护对模型精度的影响。最简做法是在服务端聚合前对每个客户端的权重做范数裁剪然后加高斯噪声。import torch def clip_and_noise_aggregate(client_updates, client_sizes, clip_norm1.0, sigma0.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_stateclip_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(seed42): # 固定 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 Falsecudnn.deterministicTrue会牺牲一点性能但换来的是卷积计算的确定性。benchmarkFalse禁止 cuDNN 在运行时自动选择算法否则算法选择的随机性也会影响结果。之后把你的实验记录落盘成 CSV每行一条import csv def log_round(row, pathfl_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, fieldnameslist(row.keys())) if f.tell() 0: writer.writeheader() writer.writerow(row)记录字段至少包括当前轮次、全局测试准确率、学习率、本地 epoch 数、客户端参与比例、Dirichlet 的 alpha、随机种子、数据切分版本号。我自己的习惯是把这些信息直接拼进文件名比如cifar10_a0.5_ratio0.5_lr0.001_r42_v3.csv这样就算日志文件堆满一个文件夹也不会搞混哪份实验对应哪组参数。这是我做联邦学习实验吃过亏之后养成的习惯——不固定种子、不落盘指标你的代码跑出来的任何“结论”都只是巧合。希望这份笔记能帮你把横向联邦图像分类的学习代码跑通并跑出能说服自己的结果。本文还有配套的精品资源点击获取