ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

横向联邦图像分类实战:从零实现FedAvg聚合与PyTorch代码详解

2026/9/12 16:02:48 拓冰建站 浏览量
横向联邦图像分类实战:从零实现FedAvg聚合与PyTorch代码详解 简介基于Python从零实现横向联邦图像分类的学习代码配套《联邦学习实战》第3章面向联邦学习入门者、计算机相关专业学生及课程设计/毕业设计人员帮助理解横向联邦学习框架下的图像分类完整流程。代码采用服务端与客户端分离结构注释充分运行环境需Python、PyTorch和CIFAR10数据集包内提供配置文件与运行说明可快速复现。整个资源共22个文件以Python脚本、JSON/XML配置、Markdown说明和示例图片为主压缩包仅156KB轻量便携。目前已有125人学习下载代码测试通过运行稳定。资源覆盖数据加载、模型定义、客户端训练、服务端聚合等模块目录结构清晰配合说明文档可帮助从零搭建横向联邦图像分类流程同时理解联邦平均等核心思想既可作为课程设计或毕业设计的参考也可作为联邦学习实践素材并支持二次扩展。1. 从零理解横向联邦图像分类为什么学它以及它和集中式学习的边界“横向联邦图像分类”这个名字初看会以为它只是把数据扔进 CNN 然后跑个 accuracy 那么简单。真正落地时你会遇到一个完全不同的难点数据分布在不同参与方手里而这些参与方彼此不能看对方的原始图片。比如医院 A 有肺部 X 光片医院 B 也有但两家机构出于隐私和合规原因不能把影像拷到同一台服务器。此时你想训练一个能同时利用两家数据的分类模型横向联邦学习就是最常见的解法——各参与方在本地用自己的数据训练只把模型参数或者梯度发给中心服务器由服务器聚合后再分发给各方如此迭代直到收敛。我推荐所有人从“学习代码”的角度切入横向联邦而不是先啃论文原因在于联邦学习的工程坑远多于算法坑。数据划分方式写错、参数聚合时把未参与客户端的权重也平均进去、训练轮次与本地 epoch 的对应关系没理清这些问题不看代码根本发现不了。这篇文章会带着你从环境搭建、数据划分、模型构造到服务端聚合流程逐步写出可运行的横向联邦图像分类代码并且每一步都保留完整注释。适合的读者是懂 Python 基础、会用 PyTorch 训练过普通图像分类模型、但从未接触过联邦学习的人——以及想在实际项目中快速搭建横向联邦基线方案的工程师。如果你是第一次接触“横向”这个词先记住一个区别横向联邦指不同参与方拥有相同特征空间、不同样本空间比如两家的图片都是 300×300 的 RGB但人不一样纵向联邦则是特征互补那套逻辑完全不同。2. 横向联邦图像分类的学习代码先把理论打牢从 FedAvg 到样本不平衡2.1 横向联邦的核心公式FedAvg 聚合与本地更新横向联邦最经典的算法是 FedAvg它的思路极其朴素每个参与方在本地数据上训练若干轮然后把模型权重发给服务器服务器对各权重做加权平均。用公式表示就是第 t 轮服务器把当前全局模型 w_t 分发给所有参与方参与方 k 用自己的数据训练得到 w_t^{k}服务器计算 w_{t1} Σ (n_k / N) · w_t^{k}其中 n_k 是参与方 k 的样本数N 是总样本数。关键在于加权系数是样本数量占比而不是参与方数量占比。很多初学者直接对所有权重做算术平均这在各参与方数据量相同时没问题但一旦某一家有 10 万张图、另一家只有 1 万张算术平均会让小数据方的主导权过大模型严重偏向大数据方。学习代码时你要把职责分清楚客户端负责本地训练和返回权重服务端只负责加权聚合训练逻辑不需要理解太复杂的密码学协议。在动手前我建议你先用下面这段伪代码建立整体心智模型。它不依赖任何框架能帮你确认自己理解的是“联邦”这个动作而不是“分布式训练”。# federated_avg_pseudo.py # 伪代码展示横向联邦的一轮完整流程注重流程而非性能 global_model CNN() # 全局模型 for round_id in range(10): # 共进行 10 轮联邦迭代 selected_clients choose_clients(all_clients, fraction0.8) weights_list [] total_samples 0 for client_id in selected_clients: local_model copy.deepcopy(global_model) # 从最新全局模型开始 local_model.train() train_loader create_dataloader(client_id) # 每个客户端只用自己的数据 for epoch in range(5): for images, labels in train_loader: loss compute_loss(local_model, images, labels) optimizer.zero_grad() loss.backward() optimizer.step() weights_list.append((client_id, local_model.state_dict(), len(train_loader.dataset))) total_samples len(train_loader.dataset) # 服务端加权平均 new_weights weighted_average(weights_list, total_samples) global_model.load_state_dict(new_weights)这段代码虽然不能直接跑但标出了横向联邦最核心的四个要素全局模型的复制、客户端独立训练、权重收集、样本加权平均。多数开源框架如 Flower、PySyft 就是把这一套流程封装成接口。从零实现时不要先去看框架源码而是先把这个伪代码转换成正向设计——你会在转换过程中理解框架为什么要设计那么多回调函数和状态管理机制。2.2 样本不平衡为什么影响聚合以及如何用代码检测在横向联邦里“样本不平衡”是比普通集中式学习更隐蔽的问题。集中式学习中你可以做欠采样或过采样但联邦环境下你不能直接看到所有数据分布只能看到每个客户端返回的样本数量。因此代码里必须显式处理两种不平衡第一是参与方之间的样本量差异。FedAvg 里用n_k / N做权重系数能缓解但如果某个客户端只有 50 张图它的本地模型早早就过拟合返回的权重会引入噪声。常见做法是为每个客户端设置最小样本数阈值少于阈值的客户端在这一轮跳过训练。第二是参与方内部的类别分布倾斜。比如客户端 A 只有猫的图片客户端 B 只有狗的图片全局模型无法同时学会两个类别。此时要引入FedProx的思路——在本地损失函数上增加一个近端项限制本地模型参数偏离全局模型不要太远。下面这段代码展示了如何在 PyTorch 中实现近端项# proximal_term.py # 在损失函数中加入 ||w_local - w_global||^2 的近端正则项 import torch def fedprox_loss(model, global_weights, images, labels, mu0.01): criterion torch.nn.CrossEntropyLoss() ce_loss criterion(model(images), labels) proximal_term 0.0 for name, param in model.named_parameters(): if name in global_weights: # 计算当前参数与全局参数的 L2 距离 proximal_term ((param - global_weights[name]) ** 2).sum() return ce_loss (mu / 2) * proximal_termmu参数的默认值建议从 0.01 开始调。mu太大模型会懒得学新数据太小则退化成普通 FedAvg。这段代码的意图不是让你一定要用 FedProx而是提醒你如果学习代码时只学 FedAvg 的爽快路径遇到真实数据分布时会不知道往哪个方向改进。先写一个能够监控每个客户端 per-class 分布的 logging 函数每轮训练后记录各客户端的类别直方图比一开始就调优算法重要得多。2.3 理论到代码的映射用表格对齐概念与变量名从零实现时最劝退的地方是变量命名混乱。你可以先用一个表格把论文里的名词映射到代码变量贴在代码文件头部作为注释这样自己和后续接手的人都不会迷路。论文概念代码变量说明Global modelglobal_model服务器维护的模型每轮更新Client modelclient_model客户端从全局模型复制后本地训练的模型Roundcommunication_round一次完整的“分发-训练-聚合”Local epochlocal_epoch客户端本地训练时遍历本地数据的次数Client fractionclient_fraction每轮参与训练的客户端比例Aggregation weightsample_ratio客户端样本数占总样本数比例Momentummomentum本地优化器参数注意不是服务端聚合的参数Learning ratelr本地优化器的学习率每轮不需要衰减一个容易被忽略的细节本地优化器的momentum在横向联邦中通常不参与聚合。因为动量项记录了本地训练的历史梯度方向直接聚合这些状态会与 FedAvg 的“全局模型只由权重决定”产生冲突。如果你的代码里optimizer torch.optim.SGD(model.parameters(), lr0.01, momentum0.9)聚合时只更新model.state_dict()optimizer.state_dict()不要传到服务端。在实现过程中你可以加一个断言来防止这种错误小说《代码整洁之道》里说“注释解释为什么而不是是什么”在联邦场景里注释还要解释哪些变量不能被聚合。3. 用 python 实现横向联邦图像分类的数据准备与模型定义以 PyTorch 为例3.1 使用 torchvision 下载并构造横向联邦数据集横向联邦学习代码的数据准备和普通图像分类最大区别是你必须模拟“数据原本就分散在各客户端”的场景。先在本地下载 CIFAR-10 作为演示数据集然后手动切割成多个客户端的数据分片。实际生产环境中数据已经在各机构本地不需要你切割但学习代码时模拟这个过程能帮你理解数据分区对模型收敛的影响。# data_prep.py # 模拟横向联邦的 IID 与非 IID 数据划分 import torch import torchvision import torchvision.transforms as transforms def load_cifar10(root./data): transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) ]) train_set torchvision.datasets.CIFAR10(rootroot, trainTrue, downloadTrue, transformtransform) test_set torchvision.datasets.CIFAR10(rootroot, trainFalse, downloadTrue, transformtransform) return train_set, test_set def split_iid(dataset, num_clients5): # IID 划分所有样本随机打乱后分成 N 份 indices list(range(len(dataset))) import random random.seed(42) random.shuffle(indices) per_client len(dataset) // num_clients client_data {} for cid in range(num_clients): start cid * per_client end (cid 1) * per_client if cid ! num_clients - 1 else len(dataset) client_data[cid] torch.utils.data.Subset(dataset, indices[start:end]) return client_data这段代码里的split_iid是最简单的随机切分方式。num_clients设置 5每份约 10000 张图。这里要注意Subset不会复制数据只保存索引列表内存开销很小。在横向联邦里客户端不必拥有全部类别数据这一点在训练前你要心里有数。更接近真实情况的是 non-IID 划分。常见的做法是按类别排序后分片让每个客户端只拥有两到三种类别的图片。用下面的函数模拟def split_noniid(dataset, num_clients5, shard_per_client2): # non-IID 划分按类别排序后切分成若干分片每个客户端取固定分片 labels [label for _, label in dataset] sorted_indices sorted(range(len(labels)), keylambda i: labels[i]) num_shards num_clients * shard_per_client shard_size len(dataset) // num_shards shards [sorted_indices[i*shard_size:(i1)*shard_size] for i in range(num_shards)] client_data {} for cid in range(num_clients): selected_shards shards[cid*shard_per_client:(cid1)*shard_per_client] client_indices [idx for shard in selected_shards for idx in shard] client_data[cid] torch.utils.data.Subset(dataset, client_indices) return client_datashard_per_client2代表每个客户端只拿两个类别的图这是标准的 non-IID 设置。你在运行横向联邦实验时应该同时跑 IID 和 non-IID 两种划分做对比否则无法看出联邦算法对数据分布的敏感性。这也是你在博客或者文档里写数据准备章节时最有说服力的一张图——横轴是通信轮数纵轴是测试准确率IID 曲线远高于 non-IID 曲线。3.2 定义可复用的图像分类 CNN 模型图像分类模型不必太复杂学习横向联邦时用一个小型 CNN 足够展示所有机制。我选择用一个三卷积层加两个全连接层的模型参数总量在 100 万以内这样在普通 CPU 上也能跑动不至于因为模型太大而把时间都耗在等待训练上。# model.py # 一个轻量级 CNN结构简单但足以学习 CIFAR-10 import torch import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self, num_classes10): super(SimpleCNN, self).__init__() self.features nn.Sequential( nn.Conv2d(3, 32, kernel_size3, padding1), # 32x32 - 32x32 nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2, stride2), # 32x32 - 16x16 nn.Conv2d(32, 64, kernel_size3, padding1), # 16x16 - 16x16 nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2, stride2), # 16x16 - 8x8 nn.Conv2d(64, 128, kernel_size3, padding1),# 8x8 - 8x8 nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2, stride2), # 8x8 - 4x4 ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(128 * 4 * 4, 256), nn.ReLU(inplaceTrue), nn.Dropout(0.5), nn.Linear(256, num_classes) ) def forward(self, x): return self.classifier(self.features(x))模型用了padding1来保持空间尺寸每次池化后尺寸减半从 32×32 变 16×16再变 8×8最后 4×4。Dropout(0.5)在全连接层前防过拟。CIFAR-10 有 10 个类别num_classes10。当你在教程中看到更复杂的 ResNet 时知道一点就够在横向联邦中模型结构必须保持全局一致。所有参与方本地模型的架构必须与全局模型完全对齐连inplaceTrue这种小细节都要一致否则state_dict的键名对不上加载时直接报错。Pytorch 中加载权重用load_state_dict如果你收到 “mismatch” 错误优先检查模型定义是否复制时改动过。3.3 本地训练函数细节注释与常见陷阱本地训练是每个客户端独立执行的逻辑它是横向联邦代码中最容易写错的部分。写这个函数时你需要关注几个点优化器只作用于本地模型、训练之前要把模型设为train()模式、权重更新完毕后要把新的state_dict传回服务端。下面给出完整实现# client.py # 参与方客户端的本地训练函数 import torch def local_train(model, dataloader, device, local_epochs5, lr0.01): model: 从服务端复制来的全局模型深拷贝 dataloader: 该客户端本地的数据加载器 local_epochs: 本地迭代数据集几遍 lr: 本地优化器学习率 返回: 训练后的模型 state_dict 以及样本数量 model.train() # 切换到训练模式启用 Dropout 与 BatchNorm 的训练行为 # 每次本地训练创建新的优化器避免旧优化器状态影响本轮 optimizer torch.optim.SGD(model.parameters(), lrlr, momentum0.9) criterion torch.nn.CrossEntropyLoss() total_loss 0.0 sample_count 0 for epoch in range(local_epochs): for images, labels in dataloader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) sample_count images.size(0) avg_loss total_loss / sample_count # 返回模型的权重拷本和样本数用于服务端聚合 return {k: v.clone() for k, v in model.state_dict().items()}, sample_count一段一段说。model.train()必须调用如果漏掉BatchNorm层的统计量不会更新Dropout也不会生效模型表现会异常。创建新优化器而不是在函数外部复用是因为联邦场景下每一轮的本地训练应该是一个独立会话上一轮的动量不应该干扰本轮。末尾的clone()很重要——如果直接把模型当前参数传出去后续模型再次训练会改变这些值但你在clone()时返回的就是一个独立副本保证聚合服务器拿到的权重是“训练结束那一刻”的权重。这个函数约 20 行嵌套循环用了两个for重点是理解 loss 的累计方式loss.item() * images.size(0)乘以批量大小是为了之后算平均损失时能按样本数加权。在学习代码时你可以顺手试试将local_epochs从 5 改成 1 或 10观察对收敛速率的影响——局部迭代越多每个通信轮内消耗的计算量越大但全局收敛所需的轮数通常会减少。这两个变量之间存在一个权衡是联邦学习中最重要的实验观察之一。4. 服务端与客户端的横向联邦协调代码通信、参数聚合与训练循环4.1 通信轮的一轮完整周期分发的顺序与深拷贝陷阱真实生产环境里服务端与客户端通过网络通信gRPC、REST、消息队列等。但在学习代码阶段我们可以用本地进程间的函数调用来模拟通信。这样做的最大好处是便于调试——各种状态可以直接打印不需要处理网络异常。# server.py # 模拟服务端逻辑分发、聚合、评估 import copy import torch def federated_learning_server(global_model, client_data_loaders, test_loader, num_rounds20, client_fraction0.8, devicecpu): global_model.to(device) for round_idx in range(num_rounds): print(f--- Round {round_idx 1} ---) # 1. 选择参与本轮训练的部分客户端横向联邦中可能部分客户端掉线或不参与 available_clients list(client_data_loaders.keys()) selected_clients select_clients(available_clients, fractionclient_fraction) # 2. 分发全局模型给选中的客户端并收集训练结果 collected_weights [] total_samples 0 for cid in selected_clients: # 关键每个客户端接收的是全局模型的深拷贝而不是引用 local_model copy.deepcopy(global_model) # 客户端本地训练 new_weights, n_samples local_train( modellocal_model, dataloaderclient_data_loaders[cid], devicedevice, local_epochs5, lr0.01 ) collected_weights.append((new_weights, n_samples)) total_samples n_samples # 3. 服务端聚合 new_global_weights federated_average(collected_weights, total_samples) global_model.load_state_dict(new_global_weights) # 4. 在测试集上评估全局模型可选项但建议每轮都做 test_acc evaluate(global_model, test_loader, device) print(fRound {round_idx 1} completed, test accuracy: {test_acc:.4f}) return global_modelcopy.deepcopy(global_model)这个深拷贝是必须要的。如果你直接写local_model global_model那所有客户端共享同一个模型对象训练时互相覆盖整个联邦过程就退化成“多个客户端轮流修改同一个模型”完全丧失联邦学习“并行独立训练、聚合平均”的本质。select_clients函数可以用random.sample实现每一轮随机挑选训练参与者。横向联邦一个优势是容忍客户端掉线所以你的服务端代码应该设计成收到的权重列表长度等于实际参与客户端数量而不是预期参与数量。4.2 服务端聚合的两种写法FedAvg 与带时间戳的快照式聚合聚合函数是整个服务端最核心的部分。FedAvg 的标准写法如下# federated_average.py # 按样本数量加权的联邦平均 def federated_average(weights_and_samples, total_samples): weights_and_samples: 列表每个元素是 (state_dict, 样本数) total_samples: 所有参与客户端的样本总数 # 以第一个模型权重结构为模板初始化聚合结果 first_weights weights_and_samples[0][0] avg_weights {key: torch.zeros_like(value) for key, value in first_weights.items()} for weights, n_samples in weights_and_samples: weight_ratio n_samples / total_samples # 计算占比 for key in avg_weights: avg_weights[key] weight_ratio * weights[key] return avg_weights这种写法把每一步的权重占比提前算好然后累加最后得到的avg_weights就是加权平均结果。注意累加时使用因为 PyTorch 张量是原地操作如果你想保留原始avg_weights的副本要提前.clone()。代码中的循环里同时遍历权重和样本数结构简洁清晰。写到这里建议你换一种聚合思路来对照“快照式聚合”。标准 FedAvg 假设一个通信轮内所有客户端同时开始训练、同时结束、同时上传。真实场景中由于网络延迟或机器性能差异有些客户端会延迟到达。如果你的服务端收到一个晚到的权重是用它更新下一轮的全局模型还是把它并入当前轮再算一次这时常见做法是“异步联邦平均”服务端每收到一个客户端的权重立即用这个小批量更新全局模型公式变为w_new w_old η · (w_client - w_old)。这种方式收敛更快但稳定性更差。# async_update.py # 异步更新每来一个客户端就更新一次全局模型 def async_global_update(global_model, client_weights, sample_ratio, lr0.1): client_weights: 单个客户端的 state_dict sample_ratio: 该客户端样本数占全部已到客户端总样本数的比例或固定值 with torch.no_grad(): for key in global_model.state_dict(): new_value global_model.state_dict()[key] lr * sample_ratio * (client_weights[key] - global_model.state_dict()[key]) global_model.state_dict()[key].copy_(new_value)这里的new_value是先用张量计算再用copy_原地覆盖。copy_是 PyTorch 中把张量值拷入另一个张量的原地方法如果直接赋值global_model.state_dict()[key] new_value会破坏state_dict的引用结构加载模型时可能报警告。这一点是很多批次代码里隐藏的坑。4.3 横向联邦的评估函数与日志系统学会看每一轮的趋势而不是只看最终准确率评估函数在逻辑上与普通图像分类测试函数没有区别但你需要额外关注评估对象是谁。横向联邦中一般把数据分为三类每个客户端的本地测试集、服务端的全局测试集、以及一个可选的“未见客户端测试集”。只汇报全局测试准确率会掩盖一个真实问题某个客户端本地数据上表现极差但整体平均值看起来好看。# evaluate.py # 评估全局模型在全局测试集上的表现并能按客户端分类汇报 def evaluate(model, dataloader, device): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in dataloader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() return correct / total评估函数里调了model.eval()这会关闭 Dropout 和 BatchNorm 的随机行为让输出确定化。data_loader可沿用全局测试集每轮评估一次。最终 blog 里可以给出一个典型输出--- Round 1 --- Round 1 completed, test accuracy: 0.2531 --- Round 2 --- Round 2 completed, test accuracy: 0.4127 ... --- Round 20 --- Round 20 completed, test accuracy: 0.7643看到准确率稳步上升说明代码逻辑大概率正确。如果准确率在某一轮后掉头向下先检查是否把通信轮内的num_rounds当成 epoch 使用导致本地训练太多次产生过拟合。日志里打印学习率、参与客户端数、本地 epoch 数量等信息这些都是后期调参必需的信息。5. 正确阅读和扩展横向联邦学习代码的四个技巧横向联邦学习代码的学习难度不在于代码本身多复杂而在于很多机制是跨进程、跨对象协作的。你单看local_train函数体会不到服务端聚合的意义单看fedavg也不知道客户端传来了什么。以下四个技巧是我在复现和改造成熟项目时最常用的方法。第一把代码中所有print替换成结构化的logging并增加轮次标识。比如在本地训练函数的第一行加logger.info(f[Client {client_id}] Epoch {epoch} started)。这样做的好处是当训练规模变大时你能通过过滤日志快速定位是哪个客户端出了问题而不是在满屏的 loss 值里翻找。第二使用 TensorBoard 或matplotlib画两张图第一张是“通信轮 vs 全局测试准确率”第二张是“各客户端本地准确率热力图”。第二张图能帮你直观定位数据不平衡问题——某个客户端在图上呈深色低准确率说明它的数据分布与全局不一致。联邦学习算法的调参方向很大程度来自这两张图。第三对聚合过程做单元测试。不要等到跑完整套训练才发现聚合函数写错。用一个只有两层的小模型和手工构造的 3 个权重 dict直接调用federated_average用assert验证计算结果与手工计算是否一致。# test_aggregation.py # 单元测试验证加权平均函数正确性 def test_federated_average(): import copy model_a SimpleCNN() model_b SimpleCNN() # 故意让两个模型权重不同便于验算 for param in model_a.parameters(): param.data.fill_(1.0) for param in model_b.parameters(): param.data.fill_(3.0) weights_a {k: v.clone() for k, v in model_a.state_dict().items()} weights_b {k: v.clone() for k, v in model_b.state_dict().items()} avg_weights federated_average( [(weights_a, 100), (weights_b, 300)], total_samples400 ) for key in avg_weights: # 期望值 0.25 * 1.0 0.75 * 3.0 2.5 assert torch.allclose(avg_weights[key], torch.full_like(avg_weights[key], 2.5)) print(test passed)测试代码的妙处在于当参与方数量、样本比改变时测试数据不用改因为期望值是通过公式算出来的。你能随时发现回归。第四也是最能提升学习深度的技巧——从读代码变成改代码。拿到一份横向联邦学习代码后先尝试以下三类小改动把local_epoch从 5 改成 1观察准确率下降的幅度把client_fraction从 1.0 改成 0.4模拟掉线场景把聚合函数从加权平均改成简单平均看看在 IID 和 non-IID 划分下分别有什么差异。这些实验做完后你记住的不再是“横向联邦怎么用”而是“横向联邦里什么会引发什么”对调参的理解会深得多。本文还有配套的精品资源点击获取