☰
单机模拟横向联邦学习:Python实现FedAvg与Non-IID数据切分
2026/9/25 1:03:53 网站建设 项目流程

简介:这份资源面向具备一定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_clients=5, seed=42): """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_clients=5, seed=42): """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_size=0.2, random_state=42 ) clients_iid = split_iid(X_train, y_train, num_clients=5) clients_noniid = split_noniid(X_train, y_train, num_clients=5) 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, minlength=10)}")

这段代码的关键参数是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, lr=0.01, batch_size=32): self.client_id = client_id self.X = torch.tensor(X, dtype=torch.float32) self.y = torch.tensor(y, dtype=torch.long) self.dataset = TensorDataset(self.X, self.y) self.loader = DataLoader(self.dataset, batch_size=batch_size, shuffle=True) self.model = SimpleMLP() self.optimizer = optim.SGD(self.model.parameters(), lr=lr) self.criterion = nn.CrossEntropyLoss() def train(self, global_state_dict, local_epochs=1): """加载全局参数,本地训练若干轮,返回更新后的 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], dtype=torch.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, rounds=20, clients_per_round=5): 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_epochs=1) 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, dtype=torch.float32) pred = global_model(X_test_t).argmax(dim=1).numpy() acc = (pred == y_test).mean() history.append(acc) print(f"Round {r+1:02d} | test acc = {acc:.4f}") return global_state, history

rounds是通信轮次,本地模拟一般 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 numpy==1.26.4 scikit-learn==1.4.2 pip install torch==2.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越多越接近集中式,但通信开销线性增长5
local_epochs1–5增大加速本地拟合,但 Non-IID 下易漂移1
lr0.001–0.05过大震荡,过小收敛慢0.01
rounds20–100决定总通信次数,看准确率 plateau30
batch_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。我一般会跑alpha=0.1、0.5、1.0 三组,对比 FedAvg 的收敛曲线,这样能直观看到异构程度对聚合的影响。

def split_dirichlet(X, y, num_clients=5, alpha=0.5, seed=42): """按 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, rounds=30) print(f"alpha={alpha} | final acc={hist[-1]:.4f} | " f"rounds to 0.85={next((i+1 for i,a in enumerate(hist) if a>=0.85), 'N/A')}")

这段代码里alpha是 Dirichlet 分布的集中度参数,rng.dirichlet([alpha] * num_clients)为每个类别生成一组客户端比例,再按比例把该类样本分给各客户端。跑完三组后重点看两个指标:最终准确率和达到 85% 准确率所需轮次。经验上alpha=0.1时客户端之间几乎不共享类别,FedAvg 可能需要 25 轮以上才能到 85%,而alpha=1.0时十几轮就够了。如果alpha=0.1下曲线一直震荡,可以试试每轮多选几个客户端参与聚合,或者把学习率降到 0.005。

从那以后我每次做联邦学习模拟,都会先把 Dirichlet 切分和 IID 切分的基线都跑一遍,确认聚合逻辑在两种分布下都正常,再往上加新策略。这个习惯帮我省掉了不少“以为是算法问题、其实是数据切分写错了”的后悔药。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询