简介这份资源面向具备一定Python与深度学习基础、希望动手理解联邦学习机制的开发者与学习者围绕「本地模拟横向联邦学习」这一主题提供一套可运行的完整工程。它解决的是在单机环境下模拟多客户端协作训练、避免真实分布式部署门槛的问题适合作为入门联邦学习、验证聚合算法与通信流程的实践素材。压缩包共25个文件约302.68MB以py脚本为核心辅以pyc缓存、xml与iml工程配置、json配置、html说明及CIFAR-10数据集分片文件涵盖服务端、客户端、模型与数据加载等模块目录结构清晰。已有723人学习下载。读者可据此理解客户端本地训练、参数上传、服务器聚合并广播全局模型的完整闭环掌握模型定义、数据划分与通信接口设计思路并在此基础上尝试差分隐私、异步更新等扩展方向。1. 从一台笔记本开始为什么横向联邦学习值得本地跑一遍很多人第一次接触联邦学习是被“数据不出域”这四个字吸引的但真到动手时又卡在“没有多台设备、没有真实边缘节点”上。其实横向联邦学习的核心逻辑——多个客户端各持一部分样本、共享同一套特征空间、在本地训练后只上传参数——完全可以在单机用 Python 模拟出来。你不需要 GPU 集群也不需要真实分布在各地的终端一台装了 Python 的笔记本就能把 FedAvg 的完整链路跑通。这份资源就是围绕“本地模拟横向联邦学习”展开的 Python 实现适合想入门联邦学习但被环境劝退的开发者、需要快速验证聚合策略的研究者以及想在自己数据上试水联邦训练的工程师。它解决的不是“生产级部署”而是“先让流程在你手里跑起来、看得见每一轮参数怎么变”。2. 横向联邦的本地模拟从数据切分到 FedAvg 聚合2.1 横向联邦的数据分区逻辑与 IID/Non-IID 切分横向联邦学习Horizontal Federated Learning的前提是各参与方拥有相同的特征维度、不同的样本集合。放到本地模拟场景里就是把一份完整数据集按样本维度切给 N 个虚拟客户端。最直接的做法是用numpy.array_split做均匀切分这样每个客户端拿到的类别分布接近全局分布也就是 IID独立同分布场景。但真实边缘设备的数据往往是非独立同分布的比如某个客户端全是数字 0 和 1 的样本另一个客户端全是 7 和 8。为了模拟这种 Non-IID 情况常见做法是按标签排序后再切分或者用 Dirichlet 分布控制每个客户端的类别比例。我一般会先写一个数据切分函数把 IID 和 Non-IID 两种模式都留出来方便后续对比聚合效果。下面这段代码用sklearn的 digits 数据集做演示它比 MNIST 轻量本地跑几十轮也不会等太久。import numpy as np from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split def split_iid(X, y, num_clients5, seed42): IID 切分随机打乱后均匀分配 rng np.random.default_rng(seed) indices rng.permutation(len(X)) splits np.array_split(indices, num_clients) return [(X[idx], y[idx]) for idx in splits] def split_noniid(X, y, num_clients5, seed42): Non-IID 切分按标签排序后分段模拟数据异构 indices np.argsort(y) splits np.array_split(indices, num_clients) return [(X[idx], y[idx]) for idx in splits] # 加载数据并归一化 digits load_digits() X digits.data / 16.0 # 归一化到 [0,1] y digits.target X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.2, random_state42 ) clients_iid split_iid(X_train, y_train, num_clients5) clients_noniid split_noniid(X_train, y_train, num_clients5) print(IID 各客户端样本数:, [len(c[1]) for c in clients_iid]) print(Non-IID 各客户端标签分布:) for i, (_, yy) in enumerate(clients_noniid): print(f client {i}: {np.bincount(yy, minlength10)})这段代码的关键参数是num_clients和seed。num_clients决定虚拟客户端的数量一般设 5 到 10 就能看出聚合趋势seed保证每次切分结果可复现调参对比时不会因为数据顺序变了导致结论漂移。Non-IID 切分里用np.argsort(y)把同标签样本聚在一起再分段这样每个客户端的标签分布会明显偏斜更接近真实场景。跑完可以打印一下各客户端的标签直方图如果发现某个客户端只有两三个类别说明 Non-IID 程度已经很高了后续聚合时全局模型可能会出现震荡。2.2 客户端本地训练用 PyTorch 写一个可复用的 LocalUpdate本地训练是联邦学习里最容易被低估的一环。很多人以为“本地训练”就是普通训练但在联邦场景下客户端模型必须和全局模型保持完全一致的结构否则上传的参数字典对不上聚合时直接报错。我习惯把客户端封装成一个类里面持有模型、优化器和本地数据对外只暴露train方法返回 state_dict 和样本数。import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset class SimpleMLP(nn.Module): 一个轻量 MLP输入 64 维digits 特征输出 10 类 def __init__(self): super().__init__() self.net nn.Sequential( nn.Linear(64, 128), nn.ReLU(), nn.Linear(128, 64), nn.ReLU(), nn.Linear(64, 10) ) def forward(self, x): return self.net(x) class FederatedClient: def __init__(self, client_id, X, y, lr0.01, batch_size32): self.client_id client_id self.X torch.tensor(X, dtypetorch.float32) self.y torch.tensor(y, dtypetorch.long) self.dataset TensorDataset(self.X, self.y) self.loader DataLoader(self.dataset, batch_sizebatch_size, shuffleTrue) self.model SimpleMLP() self.optimizer optim.SGD(self.model.parameters(), lrlr) self.criterion nn.CrossEntropyLoss() def train(self, global_state_dict, local_epochs1): 加载全局参数本地训练若干轮返回更新后的 state_dict 和样本数 self.model.load_state_dict(global_state_dict) self.model.train() for _ in range(local_epochs): for bx, by in self.loader: self.optimizer.zero_grad() loss self.criterion(self.model(bx), by) loss.backward() self.optimizer.step() return self.model.state_dict(), len(self.dataset)这里有几个参数值得注意。local_epochs控制客户端在每轮通信前本地训练多少遍设 1 就是 FedAvg 原始论文的默认设置设大了会加剧 Non-IID 下的客户端漂移。lr学习率我一般先用 0.01 跑通如果 loss 不降再调到 0.001。batch_size在本地模拟里影响不大32 或 64 都行。train方法接收global_state_dict并强制加载这一步是保证所有客户端从同一个起点出发的关键漏掉的话聚合就变成“各训各的”了。2.3 服务端聚合FedAvg 的加权平均与通信轮次控制服务端要做的事很清晰初始化全局模型每轮把参数下发给选中的客户端收齐后按样本数加权平均再进入下一轮。加权平均是 FedAvg 和简单平均的核心区别——样本多的客户端对全局模型贡献更大这在 Non-IID 场景下能缓解小客户端被淹没的问题。import copy def fedavg_aggregate(client_updates): client_updates: list of (state_dict, num_samples) total_samples sum(n for _, n in client_updates) aggregated copy.deepcopy(client_updates[0][0]) for key in aggregated.keys(): aggregated[key] torch.zeros_like(aggregated[key], dtypetorch.float32) for state_dict, n in client_updates: weight n / total_samples for key in aggregated.keys(): aggregated[key] state_dict[key].float() * weight return aggregated def run_federated(clients, global_model, rounds20, clients_per_round5): global_state global_model.state_dict() history [] for r in range(rounds): selected clients[:clients_per_round] # 本地模拟直接全选 updates [] for client in selected: state, n client.train(global_state, local_epochs1) updates.append((state, n)) global_state fedavg_aggregate(updates) # 每轮在测试集上评估全局模型 global_model.load_state_dict(global_state) global_model.eval() with torch.no_grad(): X_test_t torch.tensor(X_test, dtypetorch.float32) pred global_model(X_test_t).argmax(dim1).numpy() acc (pred y_test).mean() history.append(acc) print(fRound {r1:02d} | test acc {acc:.4f}) return global_state, historyrounds是通信轮次本地模拟一般 20 到 50 轮就能看到收敛趋势clients_per_round控制每轮参与聚合的客户端数量设成和总客户端数一样就是全参与设小一点可以模拟部分参与的场景。fedavg_aggregate里先把聚合张量初始化为零再按n / total_samples加权累加注意这里把 state_dict 转成 float32 再算避免整型参数在平均时被截断。跑起来后每轮打印测试准确率如果发现准确率来回跳大概率是 Non-IID 太极端或者学习率偏大可以先把local_epochs降回 1 再观察。3. 把模拟跑稳环境配置、超参调试与结果验证3.1 Python 环境与依赖版本避开 PyTorch 和 NumPy 的兼容坑本地模拟横向联邦学习对环境的依赖其实不重核心就是 Python、NumPy、PyTorch 和 scikit-learn。但版本搭配不对跑起来就是各种报错。我踩过的坑里最常见的是 PyTorch 2.x 和 NumPy 2.x 的兼容问题——某些旧版 PyTorch 在 NumPy 2.0 下会报module numpy has no attribute float之类的错。稳妥的做法是建一个干净的虚拟环境把版本锁死。python -m venv fed_env source fed_env/bin/activate # Windows 用 fed_env\Scripts\activate pip install numpy1.26.4 scikit-learn1.4.2 pip install torch2.2.2 --index-url https://download.pytorch.org/whl/cpu如果你用的是 VS Code 或 PyCharm记得把解释器切到这个虚拟环境否则终端里装好了、编辑器里还是旧环境跑起来照样报ModuleNotFoundError。CPU 版 PyTorch 对本地模拟完全够用digits 数据集只有 1797 个样本每轮训练几秒钟就完事没必要折腾 CUDA。装完之后用python -c import torch; print(torch.__version__)确认一下版本再跑主脚本。3.2 超参数对照实验客户端数量、本地轮次与学习率把流程跑通只是第一步真正让模拟有价值的是做对照实验。我一般会固定其他变量单独调一个参数看测试准确率曲线的变化。下面这张表是我在 digits 数据集上跑出来的经验值供你起步参考。参数常用范围对收敛的影响建议起步值num_clients5–20越多越接近集中式但通信开销线性增长5local_epochs1–5增大加速本地拟合但 Non-IID 下易漂移1lr0.001–0.05过大震荡过小收敛慢0.01rounds20–100决定总通信次数看准确率 plateau30batch_size16–128本地模拟影响小主要影响训练速度32做实验时建议把 IID 和 Non-IID 两组结果画在同一张图上。IID 下 FedAvg 通常十几轮就能到 90% 以上Non-IID 下可能要到 30 轮以后才稳定而且最终准确率会低几个百分点。如果 Non-IID 曲线剧烈震荡可以试试把local_epochs降到 1、或者把lr减半这两个操作对稳定性的提升最明显。3.3 聚合结果验证全局模型评估与参数一致性检查跑完训练不能只看最后一行准确率得确认聚合逻辑真的生效了。我习惯在每轮聚合后做两件事一是用测试集评估全局模型二是抽查聚合后的参数是否等于各客户端参数的加权平均。第二点听起来多余但如果你在fedavg_aggregate里不小心用了简单平均、或者权重算错了光看准确率不一定能发现。# 验证聚合正确性手动算一个参数的加权平均和聚合结果对比 def verify_aggregation(client_updates, aggregated): total sum(n for _, n in client_updates) key list(aggregated.keys())[0] manual sum(sd[key].float() * (n / total) for sd, n in client_updates) diff (manual - aggregated[key]).abs().max().item() print(f聚合校验 | key{key} | max diff {diff:.6f}) assert diff 1e-5, 聚合结果与手动加权平均不一致这个校验函数在调试阶段非常有用尤其是当你改了聚合策略、想确认加权逻辑没写反的时候。max diff应该接近 0如果大于 1e-5说明聚合里混入了额外操作或者权重计算有误。另外全局模型评估时记得先load_state_dict再eval()漏掉eval()的话 BatchNorm 和 Dropout 层的行为会和训练时不一致准确率会偏低。4. 避坑与排查本地模拟联邦学习最容易翻车的五个地方现象一聚合时 state_dict 的 key 对不上报Unexpected key(s) in state_dict。原因通常是客户端模型和服务端模型结构不一致比如一个用了nn.Sequential、另一个手动写了forward层名不同。解决方法是把模型定义抽到一个公共模块里客户端和服务端都从同一个类实例化不要各写各的。现象二Non-IID 下准确率一直上不去甚至比单客户端本地训练还低。这是客户端漂移的典型表现。每个客户端在本地多轮训练后模型参数已经偏向自己的数据分布加权平均后反而互相抵消。先把local_epochs设回 1如果还不行就降低学习率或者改用 FedProx 这类带近端项的聚合策略。现象三每轮准确率波动超过 5 个百分点曲线像锯齿。常见原因是客户端采样不稳定或者数据切分时没固定随机种子。检查split_iid和split_noniid里的seed是否固定以及DataLoader的shuffle是否引入了不可复现的随机性。把torch.manual_seed和np.random.seed在脚本开头都设一遍。现象四训练 loss 正常下降但测试准确率始终在 10% 左右十分类等于随机猜。大概率是标签和输出维度对不上或者数据归一化时把特征缩放到异常范围。检查SimpleMLP最后一层输出是不是 10以及X归一化后是否还在合理区间。digits 数据集原始特征范围是 0–16除以 16 归一化到 [0,1] 是常规操作漏掉这步会导致梯度爆炸或消失。现象五跑了几十轮准确率和第一轮几乎一样模型根本没更新。先确认fedavg_aggregate里是否真的把客户端参数累加进去了而不是返回了初始化的零张量。再检查client.train里是否加载了global_state_dict——如果客户端每次都用自己初始化的参数训练聚合就变成了对随机初始化的平均自然不收敛。5. 进阶技巧用 Dirichlet 分布模拟更真实的 Non-IID 并做消融对比均匀分段式的 Non-IID 还是太“整齐”了真实场景里每个客户端的类别比例往往是长尾的。用 Dirichlet 分布生成标签比例可以更细粒度地控制数据异构程度。alpha越小客户端之间的分布差异越大alpha趋近无穷时退化成 IID。我一般会跑alpha0.1、0.5、1.0 三组对比 FedAvg 的收敛曲线这样能直观看到异构程度对聚合的影响。def split_dirichlet(X, y, num_clients5, alpha0.5, seed42): 按 Dirichlet 分布给每个客户端分配类别比例 rng np.random.default_rng(seed) num_classes len(np.unique(y)) client_indices [[] for _ in range(num_clients)] for c in range(num_classes): idx_c np.where(y c)[0] rng.shuffle(idx_c) proportions rng.dirichlet([alpha] * num_clients) splits (np.cumsum(proportions) * len(idx_c)).astype(int)[:-1] for i, chunk in enumerate(np.split(idx_c, splits)): client_indices[i].extend(chunk.tolist()) return [(X[idx], y[idx]) for idx in client_indices] # 消融对比不同 alpha 下的收敛轮次 for alpha in [0.1, 0.5, 1.0]: clients [FederatedClient(i, Xc, yc) for i, (Xc, yc) in enumerate(split_dirichlet(X_train, y_train, 5, alpha))] global_model SimpleMLP() _, hist run_federated(clients, global_model, rounds30) print(falpha{alpha} | final acc{hist[-1]:.4f} | frounds to 0.85{next((i1 for i,a in enumerate(hist) if a0.85), N/A)})这段代码里alpha是 Dirichlet 分布的集中度参数rng.dirichlet([alpha] * num_clients)为每个类别生成一组客户端比例再按比例把该类样本分给各客户端。跑完三组后重点看两个指标最终准确率和达到 85% 准确率所需轮次。经验上alpha0.1时客户端之间几乎不共享类别FedAvg 可能需要 25 轮以上才能到 85%而alpha1.0时十几轮就够了。如果alpha0.1下曲线一直震荡可以试试每轮多选几个客户端参与聚合或者把学习率降到 0.005。从那以后我每次做联邦学习模拟都会先把 Dirichlet 切分和 IID 切分的基线都跑一遍确认聚合逻辑在两种分布下都正常再往上加新策略。这个习惯帮我省掉了不少“以为是算法问题、其实是数据切分写错了”的后悔药。希望帮到你。本文还有配套的精品资源点击获取