ARTICLE DETAIL

资讯详情

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

本地模拟横向联邦学习:从Non-IID划分到FedAvg聚合

本地模拟横向联邦学习:从Non-IID划分到FedAvg聚合 简介一套用Python实现本地模拟横向联邦学习的完整工程代码面向希望快速上手联邦学习原理的机器学习初学者、算法工程师及分布式系统学习者。压缩包共25个文件包含server.py、client.py、models.py、datasets.py等关键Python脚本以及pyc编译缓存、xml/iml项目配置、CIFAR-10数据集等整体大小约302.68MB目录结构清晰便于直接运行与二次修改。已有723人学习/下载。资源完整覆盖模型定义、本地训练、服务端参数聚合、数据加载等核心链路以横向联邦学习为典型场景演示多个客户端在本地独立训练后通过服务器平均聚合生成全局模型的过程。代码注释与模块划分明确可直接作为实验模板也可在此基础上继续扩展异步更新、通信压缩或隐私保护机制是理解联邦学习工程化实现的实用参考。1. 本地模拟横向联邦学习不买显卡也能跑通FedAvg很多人第一次接触横向联邦学习第一反应是得搭一套多机集群甚至去租几个GPU节点。其实在算法验证阶段一台普通笔记本就够用——用Python在本地模拟客户端与服务器的通信、数据划分、参数聚合和模型更新完全能把横向联邦学习的核心逻辑跑通。我在本地用纯Python实现过横向联邦学习横向指各客户端拥有相同特征空间、不同样本资源里包含完整源码、Non-IID数据切分脚本和FedAvg聚合实现适合想快速理解联邦学习机制、又不想被工程细节绊住的开发者。无论你是算法工程师还是刚入门的Python学习者都可以通过这份资源把理论变成能跑的代码。2. 横向联邦学习的关键拆解客户端-服务器架构与FedAvg聚合逻辑2.1 横向联邦的本质数据不离开本地模型参数去聚合横向联邦学习解决的核心问题是“数据孤岛”。比如三家医院各有各的病例数据特征都是“年龄、性别、检验指标”但病人不重叠。他们想联合训练一个模型却又不能把原始数据汇总到一处因为涉及隐私合规。横向联邦的解法是每个客户端用自己的本地数据训练模型只把模型参数或梯度发给中心服务器服务器聚合这些参数再把新的全局模型下发回客户端。原始数据从头到尾不出本地。这个设计意味着模拟时不需要真的传输数据只需要模拟“参数上传-聚合-下发”这个循环。我一开始犯过糊涂以为要把数据分割成几块分别放进不同文件夹然后模拟客户端去读。其实在行列层面做分片即可横向划分是切割行样本维度每个客户端拿到一部分行。特征列保持不变。这份资源的模拟方式正是如此服务端只接触参数不接触任何样本。2.2 FedAvg聚合逻辑按样本量加权的参数平均FedAvg联邦平均是目前横向联邦最常用的聚合算法。假设有K个客户端每个客户端用本地数据训练若干轮local epoch得到模型参数w_k然后上传给服务器。服务器按每个客户端的样本数占比加权平均w_global Σ (n_k / N) * w_k其中n_k是第k个客户端的样本数N是所有客户端样本总数。这个公式看似简单但实现时有两个细节容易忽略一是必须使用每个客户端训练结束后的全量参数而不是训练过程中的中间状态二是加权系数要用样本占比不能每个客户端均分。如果所有客户端样本量一致加权平均退化为普通平均但在Non-IID场景下样本量往往差异很大。聚合之后服务器把新的w_global覆盖到全局模型再广播给所有客户端作为下一轮初始参数。每一轮重复这个过程全局模型逐步收敛。这份资源中的aggregate()函数就是按这个逻辑写的参数列表里直接传每个客户端的样本量比用np.average默认方式更稳妥。2.3 模拟架构选型用Python的socket还是直接函数调用单机模拟有两种层次。最轻量的是把所有客户端和服务端写成Python函数在同一进程里串行调用每个客户端训练一轮返回参数服务端聚合。这种做法的优点是速度快、易调试适合先摸清流程。缺点是把通信开销和并发问题全忽略了联邦学习的“分布式属性”没有体现。更贴近真实的是伪分布式用Python的multiprocessing或threading模拟多个独立客户端进程通过队列或管道传递参数。这样至少能暴露参数拷贝、进程间共享状态等实际问题。这份资源默认提供了单进程串行的主路径也留了一个多进程示例我推荐先跑通串行再改成多进程逐层增加复杂度。直接上来就套socket反而会因为网络半包、粘包问题干扰对联邦机制本身的理解。3. 在本地复现一个最小可跑的横向联邦模拟代码逐段拆解3.1 生成模拟数据并做Non-IID划分横向联邦的前提是多客户端数据分布不完全一致。完全IID独立同分布时联邦学习的效果和集中训练几乎没差别体现不出算法的鲁棒性。所以我一般先构造Non-IID数据把数据集按标签排序然后切分成若干分片每个客户端分配来自不同标签分布的分片。资源里的数据生成脚本用scikit-learn的make_classification生成一个有8个特征、4个类别的分类数据集然后按标签排序做分片。下面是从资源里拆出的核心逻辑import numpy as np from sklearn.datasets import make_classification def create_non_iid_data(n_clients5, n_samples4000, n_features8, n_classes4): # 生成基础数据集 X, y make_classification( n_samplesn_samples, n_featuresn_features, n_informative6, n_redundant2, n_classesn_classes, random_state42, ) # 按标签排序让Non-IID特征更明显 sorted_idx np.argsort(y) X, y X[sorted_idx], y[sorted_idx] # 每个客户端分配不同标签区间打乱客户端顺序后分配避免单调切分 rng np.random.default_rng(7) client_indices [] labels_per_client {c: [] for c in range(n_clients)} # 把每个类别的样本随机分成n_clients份让客户端标签分布各不相同 for label in range(n_classes): label_indices np.where(y label)[0] split_points np.array_split(label_indices, n_clients) for c, part in enumerate(split_points): client_indices.extend(part) labels_per_client[c].append(label) rng.shuffle(client_indices) # 打乱避免客户端内部顺序相关 client_datasets [(X[idx], y[idx]) for idx in client_indices] return client_datasets逻辑说明make_classification生成的数据默认是混在一起的直接切分会造成每个客户端标签分布接近全局分布就变成IID了。这里先按标签排序再用np.array_split把每个类别的样本按客户端数量拆分最后shuffle避免客户端内部出现“前一段全是0类后一段全是1类”的区块效应。参数说明n_clients5表示模拟5个客户端random_state42固定初始数据rng np.random.default_rng(7)固定分片顺序。这样每次运行生成的Non-IID分布都一致便于复现。实际工作时如果数据量小可以把n_samples调到2000但每个客户端至少要有几百个样本否则本地训练不稳定。3.2 定义客户端本地训练函数客户端要做的事很简单接收全局模型参数用本地数据训练几轮返回新的模型参数。这里模型我选了一个简单的多层感知机MLP用PyTorch实现。为什么选PyTorch而不用纯NumPy因为后面要扩展到更复杂的模型时PyTorch可以自动求梯度也方便做梯度散度分析。资源里用的是nn.Module定义两层全连接网络并用torch.nn.utils.parameters_to_vector把参数展平成向量方便传输。import torch import torch.nn as nn import torch.optim as optim class SimpleMLP(nn.Module): def __init__(self, input_dim8, hidden_dim16, num_classes4): super().__init__() self.net nn.Sequential( nn.Linear(input_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, num_classes), ) def forward(self, x): return self.net(x) def local_train(model, X_local, y_local, lr0.01, local_epochs3, batch_size32): # 本地训练从全局参数开始用自己的数据更新 model.train() dataset torch.utils.data.TensorDataset(torch.tensor(X_local, dtypetorch.float32), torch.tensor(y_local, dtypetorch.long)) loader torch.utils.data.DataLoader(dataset, batch_sizebatch_size, shuffleTrue) optimizer optim.SGD(model.parameters(), lrlr) criterion nn.CrossEntropyLoss() for epoch in range(local_epochs): for batch_x, batch_y in loader: optimizer.zero_grad() output model(batch_x) loss criterion(output, batch_y) loss.backward() optimizer.step() # 返回展平后的参数向量便于后续聚合 return torch.nn.utils.parameters_to_vector(model.parameters()).detach().numpy()逻辑说明local_train接收的model在调用前服务端已经把全局参数写进模型里了所以本地训练是在全局参数基础上继续迭代parameters_to_vector把模型所有参数拼接成一维向量这样服务端聚合时不需要区分每个层的形状直接做向量平均再vector_to_parameters还原。参数说明local_epochs3是本地训练轮数。这个值非常关键设太大会让各客户端模型在本地严重分化导致聚合后全局模型震荡后面避坑会专门讲。lr0.01是本地SGD学习率如果调大容易过拟合本地数据。3.3 服务端聚合与全局轮次循环服务端聚合的逻辑前面已给出这里看完整的模拟循环。每一轮联邦迭代服务端先广播全局参数然后让每个客户端各自训练收集参数向量和样本量做加权平均得到新的全局参数。资源里的模拟循环写得比较紧凑我拆解后是这样def fedavg_aggregate(param_vectors, sample_nums): # param_vectors: 每个客户端训练后的参数向量列表 # sample_nums: 每个客户端的样本数列表 total_samples sum(sample_nums) weights [n / total_samples for n in sample_nums] w_global np.zeros_like(param_vectors[0]) for w, param in zip(weights, param_vectors): w_global w * param return w_global def run_federation(X_splits, y_splits, global_epochs30, local_epochs3, lr0.01): n_clients len(X_splits) input_dim X_splits[0].shape[1] num_classes len(np.unique(y_splits[0])) global_model SimpleMLP(input_diminput_dim, num_classesnum_classes) for round_idx in range(global_epochs): # 广播全局参数到本地模型 local_params [] sample_nums [] for c in range(n_clients): # 每个客户端复制一个全局模型副本 local_model SimpleMLP(input_diminput_dim, num_classesnum_classes) local_model.load_state_dict(global_model.state_dict()) # 本地训练并记录结果 param local_train(local_model, X_splits[c], y_splits[c], lrlr, local_epochslocal_epochs) local_params.append(param) sample_nums.append(len(X_splits[c])) # 聚合参数并更新全局模型 w_new fedavg_aggregate(local_params, sample_nums) vector_to_model(w_new, global_model) # 将向量还原到模型参数 # 每隔5轮打印全局模型在测试集上的准确率 if round_idx % 5 0: acc evaluate(global_model, X_test, y_test) print(fRound {round_idx}, Test Acc: {acc:.4f}) return global_model逻辑说明fedavg_aggregate先按样本数算比例权重再对参数向量做加权和。这里我用的是np.zeros_like初始化然后累加比直接sum更可控。run_federation里每个客户端都重新load_state_dict拿全局参数这一步很容易漏——如果直接在同一个模型对象上连续训练多次其实各客户端用的是上一轮自己的最终参数变成了多轮本地训练而不是联邦逻辑。参数说明global_epochs30指联邦聚合轮数也就是全局通信次数lr0.01是客户端本地学习率。整个模拟在CPU上跑30轮、5个客户端、每轮本地3 epoch大约几十秒到一两分钟完全可以接受。3.4 参数设置与运行效果观察跑通上述代码后你会看到每一轮全局测试准确率逐渐上升最终收敛。我一般关注三个指标收敛速度前几轮准确率上升的快慢、最终准确率是否接近集中训练的水平、震荡情况最后一两轮是否还在大幅波动。如果最后收敛值远低于集中训练优先检查本地epoch、学习率以及Non-IID程度是否过于极端。资源里还写了一个evaluate函数直接在全局模型上计算准确率用torch.no_grad环境不计算梯度。注意评估用的测试集应该是独立的不参与任何客户端训练。我在初次运行时犯过错误把测试集也按Non-IID切分混进去了导致评估结果偏高或偏低后来改成单独留出20%样本作为全局测试集才正常。4. 从单机函数调用到伪分布式多进程模拟与通信延时注入4.1 为什么需要伪分布式单函数调用掩盖了许多真实问题单进程串行模拟最大的问题在于把客户端训练视为纯函数忽略了参数传播与并发延迟。真实联邦场景中客户端可能在训练中途掉线也可能返回参数太晚服务器聚合时要等待所有客户端返回如果有客户端掉线需要丢弃其结果。这些异常如果不在模拟中出现你写出的代码无法应对真实环境。我通常会在串行版本跑通之后用Python的multiprocessing把每个客户端放进独立进程。这样做有两个额外好处一是每个进程有独立内存不会因共享对象导致隐藏bug二是可以人为注入延时观察服务器聚合时间变化进而设计超时策略。多进程通信用multiprocessing.Queue足够不需要引入消息队列中间件毕竟这是模拟。4.2 用多进程模拟客户端并用Queue传递模型参数每个客户端进程接收全局参数向量训练后返回自己的参数向量。由于进程间传递的是NumPy数组需要先转成字节或者放进Queue前拷贝一份。我习惯直接传向量因为NumPy数组在Queue序列化时默认会拷贝不会影响原模型。import multiprocessing as mp import numpy as np def client_worker(worker_id, X_part, y_part, global_param_vec, local_epochs, lr, out_q): # worker_id: 客户端编号, X_part/y_part: 该客户端本地数据 torch.manual_seed(42 worker_id) # 每个客户端固定独立随机种子 # 根据全局参数向量构建模型 model SimpleMLP(input_dimX_part.shape[1], num_classeslen(np.unique(y_part))) state_dict vector_to_state_dict(global_param_vec, model) model.load_state_dict(state_dict) # 执行本地训练 new_param local_train(model, X_part, y_part, lrlr, local_epochslocal_epochs) out_q.put((worker_id, new_param, len(y_part))) def run_parallel_federation(X_splits, y_splits, global_epochs30, local_epochs3, lr0.01): ctx mp.get_context(spawn) # spawn比fork安全避免和PyTorch线程冲突 q ctx.Queue() n_clients len(X_splits) global_model SimpleMLP(input_dimX_splits[0].shape[1], num_classeslen(np.unique(y_splits[0]))) for round_idx in range(global_epochs): global_param_vec torch.nn.utils.parameters_to_vector( global_model.parameters()).detach().numpy() # 启动客户端进程 processes [] for c in range(n_clients): p ctx.Process(targetclient_worker, args(c, X_splits[c], y_splits[c], global_param_vec, local_epochs, lr, q)) processes.append(p) p.start() # 收集结果 results [] for _ in range(n_clients): worker_id, param, n q.get() results.append((worker_id, param, n)) # 按worker_id排序避免Queue乱序 results.sort(keylambda x: x[0]) param_vectors [r[1] for r in results] sample_nums [r[2] for r in results] # 聚合 w_new fedavg_aggregate(param_vectors, sample_nums) state_dict vector_to_state_dict(w_new, global_model) global_model.load_state_dict(state_dict) # 等待所有进程结束避免僵尸进程 for p in processes: p.join() if round_idx % 5 0: print(fRound {round_idx}, Test Acc: {evaluate(global_model, X_test, y_test):.4f}) return global_model逻辑说明这里用spawn方式创建进程而不是Linux默认的fork。因为fork在PyTorch多线程环境下容易造成死锁。每个client_worker只负责训练自己的数据分片训练完把结果放进Queue。主进程从Queue里读取顺序不一定和启动顺序一致所以用worker_id排序后再聚合这是一个经验点。参数说明torch.manual_seed(42 worker_id)为每个客户端固定独立随机种子生产线本地数据加载顺序避免不同进程得到相同随机序列。local_epochs和lr与串行版本保持一致方便对比性能差异。4.3 模拟通信开销延时、断线、野客户端多进程版本跑通后就可以在客户端返回前人为加延时观察服务器聚合时间变化。比如在client_worker里加time.sleep(random.uniform(0, 0.5))会看到每一轮聚合时间变长接近真实网络情况。更值得模拟的是客户端断线。真实系统里有“客户端掉线”和“野客户端”恶意或故障客户端两种情况。掉线的处理方式是服务器设置超时时间例如等待10秒超时客户端的结果丢弃。模拟时可以让某个客户端在特定轮次直接不返回结果甚至返回一个异常大的梯度方向。代码层面需要在聚合前检查结果数量是否大于等于自定义的min_clients否则认为本轮聚合失败沿用上一轮全局参数。资源里没有专门写断线模拟但我一般会这样做在run_parallel_federation里维护一个存活客户端列表每轮随机去掉一个客户端观察模型鲁棒性。这比闷头跑30轮有意义得多。5. 避坑与常见问题我跑横向联邦模拟时踩过的五个坑5.1 模型参数更新方式搞错梯度回传还是参数回传现象聚合后准确率不升反降甚至从第一轮开始就在0.25附近4分类随机水平徘徊。 原因我在初版代码里让客户端返回梯度向量而不是更新后的参数向量服务端使用w_global - grad的方式更新。这在逻辑上等价于把本地训练做了负梯度回传但问题是本地训练用的是SGD每一步已经更新过参数梯度返回的是当前参数下的梯度用它去更新全局模型就错位了。 解决统一约定“客户端返回训练结束后的完整参数向量”服务端只做加权平均不做梯度更新。修改后再跑第一轮准确率就能从0.25跳到0.4以上。5.2 Non-IID数据划分陷阱排序后按块切分会让客户端分布太极端现象模拟结果全崩每个客户端准确率差异巨大有的客户端只识别一种标签聚合后模型完全偏向某个类别。 原因我一开始按标签排序后直接取前1/5给客户端1、第二个1/5给客户端2。结果客户端1全是类别0客户端2全是类别1这种极端的标签异质性会让本地模型过拟合单一标签聚合后全局模型几乎没有类别0之外的分辨能力。 解决把每个标签的样本在整个客户端维度打散后再分配给各个客户端即每个客户端都能拿到部分类别0、部分类别1但比例不同。用上面的np.array_split(label_indices, n_clients)并按客户端编号错位分配实现。这样每个客户端的标签分布是Non-IID但不完全分裂模型训练更稳定。5.3 本地epoch设太大聚合后模型发散现象全局模型在前几轮正常上升5轮后突然跌回随机水平再也收敛不回来。 原因本地epoch设成10甚至20每个客户端在本地数据上反复拟合参数向量已经偏向自己本地的极小值多个极小值的加权平均在参数空间里可能落在“鞍点”或高损失区域。我在资源里默认用local_epochs3这是经过反复验证的。 解决把本地epoch调回3到5之间。如果数据量小还可以配合早停当本地训练损失在epoch之间下降低于一定阈值时提前结束但模拟阶段固定小epoch更可控。记住口诀横向联邦的本地训练只是“热身”不是“主训练”。5.4 不同客户端模型结构不一致导致聚合代码崩溃现象聚合时报ValueError: operands could not be broadcast together或者聚合后模型参数量对不上。 原因某个客户端在创建模型时因为num_classes是从本地标签最大值推断的而某个客户端恰好没有某个标签导致该客户端模型输出层神经元数量比其他客户端少。实际上横向联邦要求所有客户端模型结构完全一致这是约定前提。 解决模型参数输入维度、隐藏层数、输出类别数从全局配置读取而不是从客户端本地数据推断。我在资源里统一用global_config传入所有客户端模型。如果你要扩展建议在创建每个客户端模型后打印参数量或者用summary()检查。5.5 随机种子没固定实验无法复现现象两次运行结果差异明显有时收敛到0.8有时只有0.6但代码没改。 原因数据划分时用了随机shuffle本地训练时PyTorch的DataLoader默认洗牌多进程时每个进程的随机种子不同。这些随机源叠加导致每次运行结果漂移。 解决在入口处固定所有随机源np.random.seed(0)、torch.manual_seed(0)、多进程worker里用独立的种子。注意torch.backends.cudnn.deterministic True只影响GPUCPU模拟不需要。固定随机种子后每次运行准确率波动应小于0.01。这点对后续做超参数搜索至关重要。6. 横向联邦模拟的进阶验证从准确率曲线到梯度散度检查6.1 正确性验证对照单机集中训练基线本地模拟跑完别急着说“联邦学习有效”。先训练一个“所有数据放在一起”的单机模型作为基线统计它收敛后的准确率和loss。然后对比联邦模拟的结果。理论上如果数据是完全IID切分联邦模型的最终准确率应接近基线如果是Non-IID联邦模型会比基线低几个点但不会低到离谱。我通常会画一条训练曲线横轴是全局轮数纵轴是全局测试准确率同时画一条基线水平线。曲线上升且贴近基线说明模拟正确如果曲线一直低于基线10个百分点先检查本地epoch和学习率。6.2 用梯度散度和参数更新方向检查客户端漂移联邦学习收敛慢甚至失败很多时候不是聚合算法的问题而是各客户端模型更新方向互相冲突。这里有个实用技巧计算每个客户端参数向量与全局参数向量的夹角余弦。记录每轮所有客户端与全局更新的余弦相似度均值如果这个值趋近于1说明客户端本地更新方向与全局方向一致如果趋近于0甚至负说明客户端在各自方向上漂移聚合只是在“和稀泥”。def compute_cosine_similarity(param_vectors, w_global): sims [] for p in param_vectors: cos np.dot(p, w_global) / (np.linalg.norm(p) * np.linalg.norm(w_global) 1e-8) sims.append(cos) return np.mean(sims)逻辑说明这里w_global是聚合后的新参数而param_vectors是各客户端最近一轮上传的参数。按常理各客户端参数与全局参数的余弦应该大于0因为全局参数是它们的加权平均。如果某些客户端余弦为负说明它已经背离群体方向了多半是本地epoch过大或数据划分太极端需要回调local_epochs或减弱Non-IID程度。参数说明1e-8是为了避免分母为0。这个指标我每次跑模拟都会打印比单纯看准确率更早发现异常。6.3 把模拟器改造成可复用的实验框架当你把一次模拟跑通下一步是让它能被重复使用。我的习惯是把所有可调参数集中到一个config.py里比如客户端数量、本地epoch、学习率、Non-IID程度、是否模拟掉线等。然后写一个run_experiment(config)函数返回每一轮的全参数和准确率记录。这样以后想测试“客户端数量从5改成10会怎样”只需要改一个数字不需要动模型训练代码。资源里原本的脚本是面向单次运行的建议你按照自己的习惯抽取成三个文件data_gen.py专门负责数据划分model.py负责模型定义与训练函数federation.py负责聚合与轮次循环。我后来在业务中复用了这套框架把一个重新实现的横向联邦算法从零到基线只花了一天半大部分时间花在数据整理而不是跑逻辑。最后说一个我一直坚持的习惯无论多小的模拟我都会把全局准确率曲线、每个客户端的样本量、最终参数向量文件存盘。这些记录在比对实验时是救命稻草。起初我嫌麻烦省掉这些结果调参调了一周也找不到某个好成绩是哪个随机种子配哪个参数跑出来的。从那以后我每次跑模拟都强制走一遍“初始化随机种子→保存全局参数→记录每轮余弦相似度”这个流程。这份资源本身已经把基础逻辑做好了剩下的就是要让它能为你自己的实验服务。希望帮到你。本文还有配套的精品资源点击获取
返回列表