横向联邦学习本地模拟:PyTorch实现FedAvg与Non-IID数据切分
发布时间:2026/9/25 1:16:02来源:尧图网络
简介面向机器学习与隐私计算入门者这是一份用Python模拟横向联邦学习的完整代码包适合希望理解分布式训练原理、又不想搭建复杂集群的开发者。资源共25个文件包括客户端、服务端、模型定义、数据集加载与主控脚本以及配置、缓存和CIFAR-10图像数据集整体约302.68MB解压后即可对照运行。已有723人学习下载。内容覆盖从本地数据划分、模型训练、参数上传到服务器端联邦平均聚合的完整流程可直接运行并观察全局模型收敛同时附带gRPC通信接口示例帮助掌握客户端与服务端交互方式。在此基础上还可继续尝试改动聚合策略、加入差分隐私或模拟异步更新深入理解真实联邦学习中的隐私保护与通信开销问题适合作为课程设计或研究实验的起点。1. 为什么我在单机上模拟横向联邦学习而不是直接上框架做横向联邦学习Horizontal Federated Learning落地时最劝退的不是算法本身而是环境Flower、FATE、TensorFlow Federated 这些框架装一遍依赖冲突就能耗掉半天而且分布式调试的复杂度会直接盖过算法本身的验证需求。所以我习惯在动手接框架之前先在本地用纯 Python 把横向联邦学习的完整链路模拟一遍——数据怎么切、客户端怎么训练、服务端怎么聚合、Non-IID非独立同分布数据分布会带来什么影响全部在单机脚本里跑清楚。这个方案特别适合两类人一类是刚接触联邦学习、想搞清楚内部机制的研究者另一类是要在真实分布式系统上做工程落地、但想先把算法正确性验证掉的工程师。它不需要 GPU不需要多机环境一台普通开发机就够了而且跑通之后你再去看 Flower 这类框架的源码会发现一切都很熟悉。2. 横向联邦学习的核心链路本地模拟到底在模拟什么2.1 横向联邦 vs 纵向联邦先搞清楚边界横向联邦学习的场景是多个参与方拥有相同特征空间、不同样本空间的数据。比如三家医院各自拥有不同患者的病历字段特征基本一致但患者群体不同。在本地模拟中我们要模拟的就是这个“数据横向切分”的过程。纵向联邦学习则是相反——相同样本、不同特征比如同一批用户在三家平台上的不同维度数据这个场景涉及样本对齐和加密求交复杂度高很多不适合作为本地模拟的第一站。我见过不少新手一上来就想模拟纵向联邦结果被样本对齐、同态加密这些前置步骤劝退。所以明确边界很重要横向联邦是联邦学习的入门场景也是本地模拟收益最大的场景。2.2 服务端-客户端架构角色划分与通信协议本地模拟横向联邦学习核心是复现一套完整的服务端Server—客户端Client交互协议。我在代码里会定义三个角色服务端负责初始化全局模型、分发模型参数、接收客户端上传的梯度或模型权重、执行聚合算法最常见的是 FedAvg联邦平均然后更新全局模型。客户端持有本地数据接收全局模型参数在本地数据上训练若干轮Epoch把更新后的模型参数或梯度上传给服务端。通信调度器在本地模拟中这个角色可以合并进服务端也可以独立出来负责控制通信轮次Communication Round。这套架构看着简单但有一些容易忽略的语义需要先立住客户端上传的到底是什么是梯度还是模型参数这两种方式在数学上等价但在实现上要注意权重衰减、正则化等细节。我建议统一上传模型参数字典state_dict逻辑清晰也不需要额外的梯度还原处理。2.3 一次完整的通信轮次从分发到聚合在本地模拟中一轮通信一个 Communication Round包含这些步骤服务端把当前全局模型的参数打包成字典格式Python 的 dict。遍历所有客户端把参数分发下去——本地模拟中就是直接调用客户端的set_params方法。每个客户端在本地数据上训练 1 到 N 个 Epoch。客户端把训练后的模型参数返回给服务端。服务端接收所有客户端的参数按 FedAvg 算法进行加权平均更新全局模型。进入下一轮。注意第三步训练 Epoch 数在本地模拟中是一个关键超参数它直接影响通信效率和数据隐私的计算方式。在真实联邦场景中客户端每轮只做少量本地训练通常 1-5 Epoch然后就要上传模型更新因为通信成本远高于计算成本。3. 用 Python 搭建本地模拟环境数据切分与客户端划分3.1 环境准备只用 NumPy 和 PyTorch 就够了本地模拟横向联邦学习不需要装任何联邦学习框架。我的推荐组合是 Python 3.10、NumPy 和 PyTorch。这三个工具已经覆盖了数据模拟、模型定义和训练逻辑。这里放一个最小环境检查脚本python -c import numpy as np; import torch; print(fNumPy {np.__version__}, PyTorch {torch.__version__})我的建议是直接用 Anaconda 创建独立环境避免污染系统 Pythonconda create -n fedlab python3.10 conda activate fedlab pip install numpy torch --index-url https://download.pytorch.org/whl/cpu参数说明这里安装了 CPU 版 PyTorch因为本地模拟的重点是验证联邦学习逻辑不是训练大模型。如果你机器上有 CUDA可以去掉--index-url参数安装 GPU 版但单机模拟一般用不上。3.2 数据切分IID 与 Non-IID 的模拟方法数据切分是本地模拟的核心技术活。它的目标是把一份完整数据集切分成 N 份模拟 N 个客户端各自持有的本地数据。切分方式有两种IIDIndependent and Identically Distributed独立同分布切分从全量数据中随机采样每个客户端拿到的数据分布和全局分布一致。实现方法是先打乱所有样本索引再按客户端数量等分。Non-IID 切分模拟真实场景中每个客户端的数据分布差异很大。比如按标签排序后切分让 Client 0 只拿标签 0-2 的样本Client 1 只拿标签 3-5 的样本这样每个客户端的数据分布就是偏斜的。我用 PyTorch 写一个数据切分工具支持两种模式import numpy as np import torch from torch.utils.data import Dataset, DataLoader, Subset def split_dataset_by_iid(dataset, num_clients, seed42): IID切分随机打乱索引均分给每个客户端 rng np.random.default_rng(seed) indices rng.permutation(len(dataset)) split_size len(dataset) // num_clients client_indices [] for i in range(num_clients): start i * split_size end start split_size if i num_clients - 1 else len(dataset) client_indices.append(indices[start:end]) return client_indices def split_dataset_by_noniid(dataset, num_clients, num_shards2, seed42): Non-IID切分按标签排序后每个客户端分配num_shards个分片 labels np.array([dataset[i][1] for i in range(len(dataset))]) sorted_indices np.argsort(labels, kindstable) num_samples_per_shard len(dataset) // (num_clients * num_shards) shards [] for i in range(0, num_clients * num_shards): start i * num_samples_per_shard end start num_samples_per_shard shards.append(sorted_indices[start:end]) rng np.random.default_rng(seed) client_indices [] for i in range(num_clients): # 每个客户端随机拿num_shards个分片 assigned_shards rng.choice(len(shards), sizenum_shards, replaceFalse) client_indices.append(np.concatenate([shards[s] for s in assigned_shards])) # 移除已分配的分片避免重复分配 remaining_shards [s for j, s in enumerate(shards) if j not in assigned_shards] shards remaining_shards return client_indices逻辑说明IID 切分的关键是rng.permutation打乱全量索引然后均分。Non-IID 切分的核心是先按标签排序再做分片Shard分配——每个分片对应某几个标签的样本客户端随机拿几个分片数据分布自然产生偏斜。num_shards参数控制客户端之间数据分布的差异程度值越小分布差异越大。这里有个容易忽略的点num_clients * num_shards必须能整除数据集大小否则最后一个分片的样本数和前面不一致会导致客户端之间本地数据量不均衡。我一般会先检查数据集大小是否能被这个乘积整除。3.3 客户端类设计本地训练的逻辑封装客户端类的职责是接收模型参数、在本地数据上训练并返回更新后的参数。这里的关键设计是状态管理——客户端需要保存自己的数据索引、模型实例和训练超参数。我写一个可复用的客户端类import torch import torch.nn as nn from torch.utils.data import DataLoader, Subset class Client: def __init__(self, client_id, dataset, indices, model_fn, devicecpu): self.client_id client_id self.local_dataset Subset(dataset, indices) self.model model_fn().to(device) self.device device def set_params(self, global_params): 服务端分发全局模型参数 self.model.load_state_dict(global_params) def train(self, epochs, lr0.01, batch_size32, criterionNone): 在本地数据上训练若干轮返回更新后的模型参数 loader DataLoader(self.local_dataset, batch_sizebatch_size, shuffleTrue) optimizer torch.optim.SGD(self.model.parameters(), lrlr) if criterion is None: criterion nn.CrossEntropyLoss() self.model.train() for epoch in range(epochs): for x_batch, y_batch in loader: x_batch, y_batch x_batch.to(self.device), y_batch.to(self.device) optimizer.zero_grad() outputs self.model(x_batch) loss criterion(outputs, y_batch) loss.backward() optimizer.step() return self.model.state_dict() def get_params(self): return self.model.state_dict()逻辑说明Subset(dataset, indices)是 PyTorch 内置的子集封装不用复制数据只保存索引视图内存友好。set_params用load_state_dict把全局参数复制到本地模型——这里是深拷贝修改本地模型不会影响服务端的参数对象。train方法返回state_dict()是一个包含全连接层权重和偏置的字典。我在实际使用中会做一个小优化训练时把self.model.train()设上但返回参数前不需要切换到eval模式因为state_dict()只取权重不涉及 Dropout 和 BatchNorm 的前向行为差异。但如果你后面要做模型验证Evaluation记得在验证前切到eval模式否则 Dropout 会导致验证结果不稳定。4. 服务端聚合与训练循环把 FedAvg 的每一行写清楚4.1 FedAvg 聚合算法数学公式到代码实现FedAvg联邦平均是横向联邦学习最经典的聚合算法。它的数学定义是w_{global}^{t1} Σ (n_k / n) * w_k^{t1}其中 n_k 是客户端 k 的本地样本数n 是所有客户端样本总数w_k 是客户端 k 的本地模型参数。含义是每个客户端的参数对全局模型的贡献度与其数据量成正比数据量大的客户端说话更有分量。实现时我用到 PyTorch 的state_dict字典结构import copy def fedavg_aggregate(client_params_list, client_sizes): 服务端FedAvg聚合 Args: client_params_list: 每个客户端的模型state_dict列表 client_sizes: 每个客户端的本地样本数列表 Returns: 聚合后的全局模型state_dict total_size sum(client_sizes) weights [size / total_size for size in client_sizes] # 初始化聚合参数为零 aggregated_params {} first_params client_params_list[0] for key in first_params.keys(): aggregated_params[key] torch.zeros_like(first_params[key]) # 加权累加每个客户端的参数 for params, weight in zip(client_params_list, weights): for key in params.keys(): aggregated_params[key] weight * params[key] return aggregated_params逻辑说明torch.zeros_like保证聚合参数的数据类型和设备与客户端参数一致。加权累加是在 CPU 上完成的即使客户端模型在 GPU 上训练聚合时也建议把参数搬到 CPU避免多 GPU 并发时的显存竞争。这里要留意state_dict的键名必须完全一致否则key循环会漏掉或报错。如果你的模型包含 BatchNorm 的running_mean和running_var键FedAvg 需要特殊处理——直接平均可能会导致全局模型统计量偏移。最简单的方式是聚合后把running_mean和running_var重置为初始值或者干脆在本地模拟中不用 BatchNorm用 LayerNorm 或 GroupNorm 替代避免这个坑。4.2 完整的训练循环从模型初始化到多轮通信下面给出完整的主训练脚本模拟 4 个客户端、通信 20 轮的横向联邦学习过程import copy from torch.utils.data import Subset, DataLoader import torch.nn as nn # 定义一个简单模型 class SimpleNet(nn.Module): def __init__(self, input_dim784, hidden_dim128, num_classes10): super().__init__() self.fc1 nn.Linear(input_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, num_classes) def forward(self, x): x x.view(x.size(0), -1) # 展平 x torch.relu(self.fc1(x)) x self.fc2(x) return x # 假设已经加载了完整数据集 dataset # 假设已经用split_dataset_by_iid生成了client_indices NUM_CLIENTS 4 COMMUNICATION_ROUNDS 20 LOCAL_EPOCHS 5 BATCH_SIZE 32 LEARNING_RATE 0.01 # 创建客户端 clients [] for client_id in range(NUM_CLIENTS): client Client( client_idclient_id, datasetdataset, indicesclient_indices[client_id], model_fnlambda: SimpleNet(), devicecpu ) clients.append(client) # 初始化全局模型 global_model SimpleNet() global_params global_model.state_dict() # 训练循环 for round_idx in range(COMMUNICATION_ROUNDS): print(f--- Communication Round {round_idx 1} ---) client_params_list [] client_sizes [] # 1. 分发全局参数各客户端本地训练 for client in clients: client.set_params(global_params) local_params client.train(epochsLOCAL_EPOCHS, lrLEARNING_RATE, batch_sizeBATCH_SIZE) client_params_list.append(local_params) client_sizes.append(len(client.local_dataset)) # 2. 服务端FedAvg聚合 global_params fedavg_aggregate(client_params_list, client_sizes) # 3. 在测试集上评估可选 if round_idx % 5 0: global_model.load_state_dict(global_params) acc evaluate_model(global_model, test_loader) print(fRound {round_idx 1} Test Accuracy: {acc:.4f})参数说明LOCAL_EPOCHS 5表示每个客户端每轮通信时在本地训练 5 个完整 Epoch。这个值偏大会增加通信轮次的计算量偏小则模型可能欠拟合。BATCH_SIZE是本地训练批次大小LEARNING_RATE用的是 SGD 的固定学习率没有用学习率调度器因为联邦学习中的学习率调度需要更谨慎的设计。关键设计点是copy.deepcopy的时机——你可能注意到了我在fedavg_aggregate中直接返回了一个新字典不会修改传入的client_params_list所以服务端的全局模型状态是独立的。但在client.set_params(global_params)这一步PyTorch 的load_state_dict是浅拷贝客户端模型和服务端模型会共享同一份权重引用客户端train()里的optimizer.step()会原地修改权重从而影响服务端的global_params。这个问题在真实场景中会通过序列化传输回避但本地模拟必须自己处理。我的做法是在fedavg_aggregate的第一步加一个深拷贝aggregated_params copy.deepcopy(client_params_list[0]) for key in aggregated_params.keys(): aggregated_params[key] torch.zeros_like(aggregated_params[key])这样确保聚合结果是全新参数对象和任何客户端参数都不共享引用。4.3 在测试集上的评估策略全局模型与本地模型的验证差异本地模拟的一个重要优势是可以同时评估全局模型和客户端本地模型这在真实分布式系统中很难做到——因为服务端拿不到测试数据集。在模拟环境中我一般在每轮通信结束时用全局模型对测试集做评估def evaluate_model(model, test_loader): model.eval() correct 0 total 0 criterion nn.CrossEntropyLoss() test_loss 0.0 with torch.no_grad(): for x_batch, y_batch in test_loader: outputs model(x_batch) loss criterion(outputs, y_batch) test_loss loss.item() * x_batch.size(0) _, predicted torch.max(outputs, 1) correct (predicted y_batch).sum().item() total y_batch.size(0) accuracy correct / total avg_loss test_loss / total return accuracy, avg_loss注意model.eval()的调用。虽然没有 Dropout 的模型在train和eval模式下前向输出一致但 BatchNorm 层的行为在这两种模式下完全不同——eval模式使用全局统计量running_mean/running_vartrain模式使用批内统计量。所以评估前必须切到eval模式评估完再切回train。在模拟环境中还有一层细腻的操作每轮通信结束时我只用聚合后的全局参数测试性能不会去测试某个客户端的本地模型。因为联邦学习的核心目标是得到一个泛化能力强的全局模型而不是每个客户端各自的本地模型。如果你对个性化联邦学习Personalized Federated Learning感兴趣那评估逻辑就完全不同了需要分别评估全局模型和本地模型。5. 本地模拟避坑Non-IID 数据分布与聚合参数的常见问题5.1 坑一Non-IID 切分不彻底客户端之间数据重叠现象客户端之间的模型差距不大但全局模型的性能始终上不去甚至比单机训练还差。刚开始我以为是聚合算法写错了后来检查数据切分才发现问题。原因split_dataset_by_noniid函数里用了rng.choice从剩余分片中选num_shards个但没有检查选中的分片之间是否有重叠——在分片数量少、num_shards大的情况下同一个分片被分配给多个客户端的概率不为零。这让客户端之间数据分布相似度升高模拟的 Non-IID 程度被稀释了。解决在分配分片后立即从剩余列表中移除已选分片代码中已经有这一步并且在分配前检查len(shards) num_clients * num_shards。更稳妥的做法是先把所有分片编号打乱再按顺序依次发给客户端保证天然无重叠rng.shuffle(shards) for i in range(num_clients): client_indices.append(np.concatenate(shards[i*num_shards:(i1)*num_shards]))5.2 坑二客户端本地 Epoch 数过大模型发散现象通信轮次进行到第 5 轮左右全局模型的损失突然飙升到无穷大测试准确率掉到接近随机猜测。原因LOCAL_EPOCHS设置的太大比如 20每个客户端在自己的本地小数据集上反复训练模型被过拟合到局部数据上。多个客户端把各自严重偏斜的模型参数上传服务端聚合后的全局模型自然崩了。解决把LOCAL_EPOCHS降到 1-5同时调低学习率。我常用的组合是LOCAL_EPOCHS3LEARNING_RATE0.01基于 MNIST 或者 Fashion-MNIST 这类简单数据集表现稳定。如果模型仍然发散用torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)对梯度做裁剪。5.3 坑三客户端数量变化时聚合权重失真现象本来 4 个客户端每轮通信时正好有 1 个客户端掉线未参与聚合。然后我直接用fedavg_aggregate聚合剩余 3 个客户端结果全局模型性能下降明显。原因fedavg_aggregate函数用client_sizes计算权重时分母用的是当前参与聚合的客户端样本数之和。但样本数少的客户端在参与轮次不多的情况下它的数据在全局分布中的比例会被低估导致聚合结果偏斜。解决维护一个全局样本数向量每个客户端的权重始终以其全局样本数参与计算。真实联邦系统中还会考虑客户端掉线率用指数移动平均来平滑样本比例weight client_sizes[k] / global_total_size # global_total_size 所有客户端样本数之和5.4 坑四忽略模型初始化的一致性现象第一次跑模拟效果很好第二次重跑同样的代码结果差异很大。原因SimpleNet()在每次创建时调用 PyTorch 默认的随机初始化Kaiming uniform / Xavier uniform导致不同客户端初始模型不同。在联邦学习中服务端分发全局参数前客户端的初始状态必须一致否则第一轮聚合就引入了随机偏差。解决在创建客户端之前先初始化一个全局模型获取它的state_dict作为种子所有客户端初始化后立即load_state_dict(global_params)global_model SimpleNet() # 先固定全局模型 global_params copy.deepcopy(global_model.state_dict()) for client in clients: client.set_params(global_params) # 在本地训练前覆盖随机初始化另外建议固定所有随机种子包括 PyTorch 和 NumPy 的种子import random random.seed(42) np.random.seed(42) torch.manual_seed(42)5.5 坑五内存占用爆炸——Subset 的隐藏开销现象模拟 20 个客户端每个客户端持有 5000 个样本代码运行到第二轮通信时内存占用飙升到 8GB。原因Subset(dataset, indices)虽然不复制数据但每个DataLoader迭代时都会生成一个新的批次数据。如果数据集本身很大比如图片数据每个客户端同时持有多个DataLoader引用和模型优化器状态内存自然不够。另一个隐藏开销是state_dict字典在通信轮次中没有被垃圾回收Python 的引用计数在这个场景下表现不佳。解决在通信轮次之间显式释放不再需要的变量del client_params_list等用gc.collect()手动触发垃圾回收。如果数据量实在太大就把数据预先转为 NumPy 数组或者 PyTorch Tensor 保存在内存中按客户端的样本索引做切分避免Subset的迭代开销。6. 进阶记录通信轮次的完整实验日志对比不同聚合策略完成基础模拟后你的下一步是建立一个可复现的实验框架。这里有几个进阶方向6.1 实验日志与结果可视化我会在每次模拟运行时记录完整日志包括每轮通信的全局模型准确率、客户端上传参数的平均范数、每轮耗时等。日志格式用 CSV 最简单但这不够完整——我更喜欢用 JSON Lines 格式因为模型参数之外的附加信息如数据分布统计也可以塞进去import json import time # 每轮通信结束时记录 experiment_log { round: round_idx, test_accuracy: acc, train_loss: train_loss, timestamp: time.time(), num_participating_clients: len(clients), avg_update_norm: avg_norm, data_distribution: {client.client_id: len(client.local_dataset) for client in clients} } with open(flogs/exp_{exp_id}_round_{round_idx}.json, w) as f: json.dump(experiment_log, f)用matplotlib画全局模型准确率曲线和损失曲线对比不同数据切分策略IID vs Non-IID下的收敛速度差异你会直观看到 Non-IID 数据分布对联邦学习的挑战。6.2 对比不同聚合算法的收敛性与鲁棒性FedAvg 是最基础的方案但它假设所有客户端数据量对全局模型的贡献均匀。我在模拟中会额外实现 FedProx添加近端项约束客户端本地更新幅度和 SCAFFOLD修正客户端更新方向的偏移对比它们在不同 Non-IID 程度下的效果。FedProx 的实现比 FedAvg 复杂不了多少——只需要在客户端本地训练时给损失函数加一个正则项# 客户端本地训练时FedProx 的损失函数 proximal_term 0.0 for local_param, global_param in zip(model.parameters(), global_params_list): proximal_term (local_param - global_param).norm(2) ** 2 loss criterion(outputs, y_batch) (mu / 2) * proximal_term其中mu是近端项系数global_params_list是服务端分发的全局参数列表。这个改进让客户端在本地训练时不能过分偏离全局模型对 Non-IID 场景尤其重要。SCAFFOLD 则更复杂需要服务端和客户端各自维护一个控制变量Control Variate修正客户端更新方向。我建议先实现 FedProx跑通了再挑战 SCAFFOLD。6.3 从模拟到真实的最后一公里我踩过的坑本地模拟和真实分布式系统的差距主要在通信延迟和异步性。在本地模拟中所有客户端在同一台机器上训练通信是模拟的直接函数调用零延迟所以 FedAvg 天然是同步的。真实系统中客户端计算能力不均有的客户端训练快有的慢同步聚合会等待最慢的客户端——这就是异步联邦学习的动机。我的建议是先在本地模拟中把同步逻辑跑通再引入模拟延迟每个客户端sleep随机时间来测试异步聚合策略。这一层做准备后接入 Flower 框架做真实的分布式部署整体的切换成本会低很多。另外有一个血泪经验如果你的目标是跑通一个能发表论文的实验对比比如 FedAvg vs FedProx 在 CIFAR-10 上的效果不要只跑一次随机种子就下结论。建议至少跑 3 个随机种子取平均值和标准差画误差条否则 Non-IID 切分的随机性会让结论反复横跳。这个方向技术上并不难难的是耐心把每个环节的假设都想清楚数据怎么切、参数怎么聚合、日志怎么记。把这套本地模拟跑通一次你对横向联邦学习的理解会比看十篇论文都深刻。希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网