ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

从零实现横向联邦图像分类:PyTorch+Socket手写FedAvg聚合

从零实现横向联邦图像分类:PyTorch+Socket手写FedAvg聚合 简介这份基于Python从零实现横向联邦图像分类的学习代码面向想要入门联邦学习与图像分类的开发者、学生及教师。资源以CIFAR10数据集为案例完整实现了横向联邦学习中客户端本地训练、服务端参数聚合与模型更新的闭环流程每个模块均附加大量中文注释清晰标注了数据加载、模型定义、训练循环、联邦通信等关键逻辑特别适合作为深度学习和隐私计算的入门教程。整套资源共22个文件包含6个Python源码、配置与工程文件、README说明文档、示例图片等压缩包仅156KB体量轻巧、结构清晰便于下载后快速运行和二次修改。代码已经测试通过可以稳定运行既能用于课程设计、毕业设计也适合自学实践目前已有125人学习下载是一款性价比很高的入门范例。1. 横向联邦图像分类到底在做什么好。假设你在一家做 AI 的公司里模型要落地但业务方两手一摊「数据不能出部门更不能进你的训练集群」。横向联邦Horizontal Federated Learning解决的就是这个问题多个参与方各自持有同一批特征空间、不同样本的数据在数据不出本地的约束下协作训练同一个图像分类模型。区别在于单机训练里数据loader从磁盘读图而横向联邦里每个客户端用自己的数据训练只把模型权重发出去。这个标题里真正值钱的不是「图像分类」——它用到的模型就是普通的 CNN 或 ResNet——而是两件事横向两个字对应的样本划分协议以及从零实现时你必须自己写的那套聚合逻辑。市面上的 FL 框架PyTorch 自带的也不多会帮你把聚合、通信、采样都封装掉坏处是你调完参还是说不清「FedAvg 加权时到底加了什么权」。所以这篇博客的做法是自己写 socket 通信不引额外联邦框架。目标读者是会用 PyTorch 做分类、但没有摸过联邦学习的工程师或研究生。你不需要分布式训练背景只需要懂 Python、torch、socket并且愿意把代码逐行读下去。2. 先复现横向联邦的数据划分非 IID 下的 MNIST 训练集2.1 横向联邦的数据形态同样特征不同样本横向联邦里「横向」的含义按样本维度切分参与方 A 有 id 15000 的样本参与方 B 有 id 500110000 的样本所有参与方的字段定义一致。对应到图像任务就是每张图的通道数、尺寸、标签空间都相同只是图像本身不同。这一点决定了你后续做模型聚合时网络结构必须是同一个——你在客户端跑 ResNet18服务端初始化的也必须是 ResNet18。初学者最容易误把横向联邦当成「数据并行、梯度同步」的分布式训练。两者外形很像但有个本质差异数据并行的 worker 由同一个中心节点管理可以随时拿全量梯度做 AllReduce横向联邦的客户端是独立的进程/设备可能跑在不同机器上服务端拿不到梯度能拿到的只是「每个客户端训练完之后的模型权重」。所以横向联邦的起点不是通信而是「你打算给每个客户端发什么样的数据」。2.2 用 torchvision 构造多客户端数据集直切与分片对比为了复现一个「横向」场景最直接的做法是把 MNIST 训练集按样本下标切成 N 份每份交给一个客户端。这种切法得到的每个客户端数据分布与全局分布一致称为 IID独立同分布。但真实生产里各客户端的数据往往是偏斜的——某台手机上的照片几乎全是风景另一台全是宠物这种分布叫 Non-IID。下面这段代码用最简单的分片法生成 5 个客户端的数据每个客户端 12000 张图片import torch from torch.utils.data import Subset, DataLoader from torchvision import datasets, transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) full_train datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) num_clients 5 per_client len(full_train) // num_clients # 60000 // 5 12000 client_loaders [] for i in range(num_clients): indices list(range(i * per_client, (i 1) * per_client)) subset Subset(full_train, indices) loader DataLoader(subset, batch_size64, shuffleTrue, num_workers2) client_loaders.append(loader) for i, loader in enumerate(client_loaders): counts torch.zeros(10, dtypetorch.int32) for _, labels in loader: counts.scatter_add_(0, labels, torch.ones_like(labels)) print(fclient {i}: {counts.tolist()})逻辑说明full_train下载的是完整 MNIST训练集 60000 张Subset只保留指定下标DataLoader负责按 batch 产出。分片后每个客户端样本数相同但因为 MNIST 本身按标签排序存储直接顺序切分会让前几个客户端几乎只看到若干个数字——正好模拟一种偏斜的 Non-IID。如果你要 IID需要在切分前对indices做random.shuffle。参数说明num_clients5是逻辑客户端数量实际运行时一个客户端就是一个独立进程batch_size64是每个客户端本地的 batch 大小这个值受客户端显存约束服务端聚合不关心它num_workers2是 loader 子进程数Windows 下记得放进if __name__ __main__里否则多进程会重复执行下载逻辑。上面打印的counts就是标签分布如果有的客户端缺少某一类后面准确率会明显分化——这正是横向联邦要抗住的场景。2.3 用 Dirichlet 分布模拟更真实的 Non-IID直接顺序分片只能产生「某个客户端缺某几类」的硬偏斜真实世界更多是「比例偏斜」。更常用的做法是让每个客户端从各类别里按 Dirichlet 分布抽样用参数alpha控制偏斜程度。alpha越小分布越极端。以下是常见做法import numpy as np rng np.random.default_rng(42) n_classes 10 alpha 0.5 # 偏斜系数越小越偏斜 client_proportions rng.dirichlet([alpha] * num_clients, sizen_classes) # shape: (10, 5)第 j 行表示标签 j 分配给各个客户端的比例逻辑说明rng.dirichlet([alpha] * num_clients, sizen_classes)生成 10 行、5 列的矩阵第 j 行之和为 1代表第 j 类样本按什么比例分发到 5 个客户端。拿到这个比例矩阵后按np.random.choice给每个样本指派客户端即可。alpha0.5时多数类别的数据会集中在少数几个客户端alpha100时比例趋于均匀。提示做横向联邦实验时建议同时保留一个 IID 划分作为对照组。IID 下的收敛速度与最终精度是判断你代码逻辑是否正确的基线如果 Non-IID 没收敛而 IID 也没收敛大概率不是数据划分的问题而是通信或聚合代码有 bug。3. 把图像分类模型做成横向联邦客户端训练、序列化与发送3.1 客户端只做三件事拿权重、本地训练、回传权重横向联邦客户端在每一轮要执行的操作可以简化成循环从服务端拿到当前全局模型权重用本地数据训练若干 epoch然后把训练后的权重回传。注意这里回传的既不是梯度也不是 loss而是state_dict——因为服务端不需要关心每个客户端用什么优化器、跑了几步它只做加权平均。如果你在通信层传梯度那么客户端之间的迭代次数不一致时梯度的尺度就完全不可比聚合会出问题。同时还要约定哪个模型在网络两端共享。我在实现里让服务端持有模型定义一个三层的FedCNN客户端不自己初始化权重——客户端第一次连接时由服务端下发全局模型参数。这样最稳妥避免了「不同客户端初始化方式不同导致模型结构漂移」的坑。3.2 模型定义一个足够说明问题的小型 CNN图像分类任务在 MNIST 上不需要上 ResNet用两层卷积加全连接就能跑到 99% 以上做联邦学习实验也更容易观察收敛。重点在于state_dict的键名在服务端和客户端必须完全一致否则load_state_dict直接报错。import torch.nn as nn class FedCNN(nn.Module): def __init__(self, num_classes10): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(inplaceTrue), nn.Linear(128, num_classes), ) def forward(self, x): return self.classifier(self.features(x))逻辑说明输入是 1×28×28 的单通道图经过两次「卷积池化」后空间尺寸从 28×28 降到 7×7通道变成 64展平后是 64×7×73136 维再过两层全连接输出 10 个类别 logits。这个网络结构在单机 MNIST 上轻易过 99%在联邦场景下是理想基线。padding1保证卷积不改变空间尺寸这样后面的全连接输入维度是可预先算定的。注意不要在客户端模型里加 BatchNorm 的track_running_statsFalse除非你明确知道自己在干什么。联邦学习里每个客户端本地 BatchNorm 统计量只基于本地数据在 Non-IID 下会造成统计量偏移最终聚合后模型在测试集上表现异常。MNIST 这个小模型用ReLU加池化就能达到目标精度暂时不必引入 BatchNorm。3.3 客户端训练循环本地迭代与全局参数加载客户端的训练逻辑和单机训练几乎一样区别只在开头和结尾。下面这个函数独立成一个脚本client.py里的核心连接服务端、接收全局模型、训练local_epochs轮、发回新模型。import copy, pickle, socket, struct import torch def train_round(model, loader, local_epochs, lr, device): model.train() optimizer torch.optim.SGD(model.parameters(), lrlr, momentum0.9) criterion torch.nn.CrossEntropyLoss() for _ in range(local_epochs): for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() return model.state_dict() def recv_exact(sock, n): data b while len(data) n: chunk sock.recv(n - len(data)) if not chunk: raise ConnectionError(connection closed while receiving) data chunk return data def client_main(server_host, server_port, loader, devicecpu, local_epochs2, lr0.01): model FedCNN().to(device) with socket.create_connection((server_host, server_port)) as sock: # 1) 向服务端注册自己 sock.sendall(bREG) sock.recv(1) # 服务端返回 OK # 2) 接收当前全局模型 (n,) struct.unpack(!I, recv_exact(sock, 4)) global_bytes recv_exact(sock, n) global_state pickle.loads(global_bytes) model.load_state_dict(global_state) # 3) 本地训练 local_state train_round(model, loader, local_epochs, lr, device) # 4) 发送新权重 payload pickle.dumps(local_state) sock.sendall(struct.pack(!I, len(payload)) payload)逻辑说明recv_exact解决 TCP 粘包问题——sock.recv(n)不保证一次收满 n 字节所以必须循环接收直到拼接完整。struct.pack(!I, len(payload))把长度以 4 字节无符号整数摆在报文头部服务端先读长度再读报文这是最简单的长度前缀协议。pickle序列化state_dict会保留全部模型参数包括 BatchNorm 的缓冲区和全连接权重字符串键名在 pickle 里原样保留因此跨进程加载没有问题。参数说明local_epochs2是每个客户端每轮全局迭代内的本地训练轮数它直接改变单轮的通信频率——epoch 越大一轮通信内本地算得越多通信开销占比越低但过大容易让模型在本地过拟合、损害全局聚合效果lr0.01对应 SGD 的初始学习率联邦场景下常见做法是比同结构单机模型略大一些因为每轮只训练少量数据会产生近似带噪梯度的效果。待客户端大体跑通之后可以再做网格搜索local_epochs固定为 1先调lr再调num_clients与alpha。4. 服务端聚合与 FedAvg 加权你要写的核心只有十行4.1 FedAvg 为什么要按样本数加权服务端拿到 N 个客户端发回的state_dict后要做的是聚合成下一轮全局模型。最常用的算法是 Federated AveragingFedAvg公式写出来很朴素w_{t1} Σ (n_k / n) * w_{t1}^k其中 n_k 是客户端 k 的本地样本数n 是全体参与客户端样本总数。这个加权系数的含义是数据多的客户端真实样本能代表更大范围的分布它的权重在全局模型里应占更大比例。如果忽略样本数做简单平均那么数据量少的客户端会获得与数据量大的客户端同等的话语权最终模型在整体分布上会系统性偏向小数据方。在 MNIST 直切场景里各客户端样本数相等简单平均和加权平均没有区别但一旦换成 Non-IID 数据划分或客户端掉线样本数权重的价值立刻体现出来。4.2 先跑通单机版聚合不碰网络也能验证在写 socket 服务端之前我建议先写一个不依赖网络的聚合函数——用一个列表收集多个客户端训练后的state_dict直接做 FedAvg验证模型能收敛。这样把「聚合算法问题」和「网络通信问题」分开排查否则两端同时报错时很难定位。import copy def fedavg_aggregate(state_dicts, sample_nums): total_n sum(sample_nums) agg_state copy.deepcopy(state_dicts[0]) # 先把第一份拿来做初始加权的容器 for k in agg_state.keys(): agg_state[k] state_dicts[0][k] * (sample_nums[0] / total_n) for sd, n_k in zip(state_dicts[1:], sample_nums[1:]): weight n_k / total_n for k in sd.keys(): agg_state[k] sd[k] * weight return agg_state逻辑说明这里遍历每一层的参数张量把各客户端权重乘上各自样本占比再累加。关键坑有两个一是必须用copy.deepcopy(state_dicts[0])新建张量否则原地操作会污染第一个客户端的返回值二是agg_state[k] ...的张量是在 CPU 上操作的即使客户端模型在 GPU 上你也要在聚合前调用.cpu()否则浮点累加发生在不同设备上会直接报错。每个客户端的torch.Tensor可以待在原进程里一直训练聚合时只处理序列化后的字节流这样内存占用是可控的。4.3 多客户端 socket 服务端注册、收模型、聚合、广播网络版服务端用ThreadingTCPServer或裸socketthreading都可以关键约束是「一次全局轮次必须等待所有客户端到齐」。最常见的实现是服务端阻塞等待客户端注册收到REG后返回确认然后开启多线程接收各客户端的参数。等全部到齐后执行fedavg_aggregate再把聚合结果广播给所有客户端。import socket, threading, struct, pickle, time class FedServer: def __init__(self, host127.0.0.1, port6000, min_clients3): self.host, self.port host, port self.min_clients min_clients self.global_model FedCNN() self.round_idx 0 def handle_client(self, conn, client_id): try: state_bytes self._recv_len_prefixed(conn) state_dict pickle.loads(state_bytes) local_samples state_dict.pop(__num_samples__, 12000) with self.lock: self.client_states[client_id] (state_dict, local_samples) self.ready_count 1 finally: pass逻辑说明state_dict.pop(__num_samples__, 12000)使用了一个简单的协议扩展——客户端在发回权重时额外塞进一个字段记录本地样本数服务端聚合时读取并把它弹出去剩下的才是真正的模型参数。这样不需要额外的消息类型扩展通信协议时也更加平滑。注意socket 通信是同步阻塞的客户端如果中途死掉服务端会一直等它发完。实际项目里要加超时机制——conn.settimeout(60)是常见做法到了 60 秒还没有数据就丢弃该客户端并把它从本轮参与列表里移除。联邦学习对客户端掉线的容忍度来自一个默认假设每轮只需要min_clients个客户端参与不必等所有注册过的客户端都到齐。4.4 两轮聚合之间的监督指标服务端聚合完成后要验证本轮全局模型是否变好。由于训练数据不能上收验证方式是在服务端维护一份公共测试集——MNIST 的测试集只有 10000 张在这类入门实验里完全可以作为「公共测试集」使用因为横向联邦的实验目标是验证聚合算法是否工作而不是评估现实中的隐私预算。def evaluate(model, test_loader, devicecpu): model.eval() correct, total 0, 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) preds model(images).argmax(dim1) correct (preds labels).sum().item() total labels.size(0) return correct / total这段代码很常规但注意一个陷阱评估时的model必须是从agg_state重新实例化后加载的模型不能直接用服务端持有的self.global_model做 in-place 更新然后评估——如果服务端模型网络结构里含有 Dropout 或 BatchNormin-place 更新后统计缓冲区会残留上一轮的状态。安全做法是每次evaluate前都copy.deepcopy一份全局状态并load_state_dict到新模型里。5. 收敛慢、精度上不去时的可视化查因与断点续训5.1 先看曲线loss 震荡是横向联邦的常态而非 bug很多第一次做联邦实验的人看到训练 loss 曲线比单机训练抖得多就怀疑聚合写错了。实际上这是横向联邦的典型特征每个客户端每轮只见到全局数据的一小部分且 Non-IID 划分下各客户端数据分布差异大导致本地优化方向彼此不一致。判断代码有没有问题的正确做法是同时画两条曲线全局模型在公共测试集上的准确率以及每个客户端本地训练后的验证 loss。准确率曲线如果整体趋势向上而单轮有波动说明聚合正常如果准确率在前几轮完全不涨再检查数据划分是不是偏斜到某个客户端只剩一个类别。用matplotlib在服务端进程里画图每次聚合后追加一个点import matplotlib.pyplot as plt def plot_training_curve(rounds, accs, client_loss_curves, save_pathfed_curve.png): fig, ax1 plt.subplots() ax1.plot(rounds, accs, b-o, labelglobal test acc) ax1.set_xlabel(communication round) ax1.set_ylabel(accuracy, colorb) ax2 ax1.twinx() for cid, losses in client_loss_curves.items(): ax2.plot(rounds, losses, --, alpha0.5, labelfclient {cid} loss) ax2.set_ylabel(local loss, colorr) fig.legend() fig.savefig(save_path) plt.close(fig)逻辑说明twinx()建立双 y 轴左边是测试集准确率右边是各客户端 loss。这张图能直接看出「准确率上升但某个客户端 loss 不降」或「所有客户端 loss 都降但准确率不涨」。后者通常意味着模型过度拟合了某个客户端的本地分布解决办法是调小local_epochs或调大lr并配合学习率衰减。5.2 防崩溃断点续训与幂等握手横向联邦实验最烦的不是模型不收敛而是训练到第 13 轮时候服务端崩了全部白跑。因为客户端有 5 个、每个本地训练跑 2 个 epoch一轮就要几十秒跑几十轮需要几十分钟。断点续训的正确姿态是服务端每轮聚合完成后立即把agg_state、round_idx、当前最优准确率打包保存成checkpoint.pt下次启动时检测到文件就加载。def save_checkpoint(round_idx, agg_state, best_acc, pathserver_ckpt.pt): torch.save({ round_idx: round_idx, agg_state: agg_state, best_acc: best_acc, }, path) def load_checkpoint(pathserver_ckpt.pt): ckpt torch.load(path, weights_onlyFalse) model FedCNN() model.load_state_dict(ckpt[agg_state]) return ckpt[round_idx] 1, model, ckpt[best_acc]逻辑说明round_idx保存的是「已完成的轮次」加载后下一轮的全局轮次要加一否则重启后第一轮会重新训练第 1 轮best_acc用来记录历史最优准确率只在刷新最优时才覆盖保存这样即使后面几轮效果变差你仍能找回最优模型做推理。配合服务端重启后重新监听端口、客户端在断线后自动重连整套实验挂在后台跑一晚上就很稳妥。5.3 一个立竿见影的改法把本地测试集独立出来最后一个调优技巧在客户端进程里训练集只用于反向传播但每轮训练完成后单独留出几百张本地数据不参与训练作为「本地验证集」。这样你能在客户端侧看到更细粒度的指标——本地验证准确率与服务端全局准确率之差可以量化 Non-IID 的偏斜程度。实现上只要在数据划分处多切一块train_idx indices[:11000] val_idx indices[11000:] client_val_loader DataLoader(Subset(full_train, val_idx), batch_size64)每轮本地训练结束后在client_val_loader上评估一次。你会发现某个客户端本地验证准确率可能早就到 99%但它回传的模型反而拉低了全局准确率——这说明你需要加大每轮参与聚合的客户端数量让全局模型不会被单个本地过拟合的模型带偏。横向联邦实验做到这一步就已经超过了大多数只跑通 demo 的教程你手上有的是数据分布、聚合权重和收敛曲线三者之间的定量关系而不是一句「效果不错」。本文还有配套的精品资源点击获取
返回列表