ARTICLE DETAIL

资讯详情

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

横向联邦学习FedAvg本地模拟:从原理到Python代码实战

横向联邦学习FedAvg本地模拟:从原理到Python代码实战 简介面向Python开发者与联邦学习入门者的本地模拟横向联邦学习项目围绕数据不出本地、多客户端协作训练的核心思想提供了一套可运行的最小实现。压缩包约302.68MB共25个文件包含server.py、client.py、models.py、datasets.py、main.py等5个Python源文件以及对应pyc缓存、PyCharm项目配置、CIFAR-10分批数据文件和训练配置JSON代码结构与数据划分清晰便于直接启动实验。已有723人学习下载。通过这套代码可直观看到客户端在本地数据上训练模型、上传参数服务器完成聚合后广播更新的完整过程项目预留了数据集加载、模型定义、训练配置等模块读者可替换为自己的网络结构或调整联邦轮次用于对比不同聚合策略与数据分布下的效果。适合作为课程设计、论文复现或联邦学习技术预研的起点。1. 横向联邦学习为什么要先本地模拟不花一分钱验证 FedAvg 的收敛性横向联邦学习Horizontal Federated Learning是多家机构在数据不出域的前提下用各自样本共同训练一个模型的方式。“本地模拟”就是在一台机器上把一份数据集按样本 ID 切成多份假装它们是不同参与方的私有数据再把 FedAvg 这类聚合算法完整跑一遍用几十分钟验证“联邦训练到底能不能收敛、效果离集中训练差多少”。很多人一上来就接联邦学习框架结果被通信、加密、平台适配搞晕连算法本身的收敛性都没验证过。我建议新项目先做一次本地模拟半天时间把算法、数据和参数边界摸清楚再谈上不上生产。这套方法适合算法工程师、隐私计算方向的研发也适合想入门联邦学习的 Python 开发者。2. 横向联邦学习原理与 FedAvg先搞懂聚合公式再写代码2.1 横向的划分所有参与方拥有同样的特征空间横向联邦是相对于纵向联邦而言的。横向意味着参与方的特征字段一致医院 A 和医院 B 都有“年龄、血压、病史”这些列但患者 ID 完全不同银行 A 和银行 B 都有“收入、负债、历史交易”这些列但客户 ID 完全不同。训练时每个参与方在本地用自己的样本训练只有模型参数或梯度离开本地方可被服务器聚合。本地模拟的数据切分就要模仿这个“同样特征、不同 ID”的格局按行把样本切分成多份而不是按列切特征。这个选择直接决定代码里切分函数的写法。如果把特征切开那就变成纵向联邦了聚合逻辑完全不同。很多新手在这里第一脚就踩歪后面所有实验结论都是错的。2.2 FedAvg 聚合公式不是简单平均是按样本量加权FedAvg 是联邦学习最经典的聚合算法核心更新公式可以写成w_{t1} sum_k (n_k / n) * w_k其中 k 是客户端编号n_k 是第 k 个客户端本地的样本数n 是所有参与客户端的总样本数w_k 是第 k 个客户端本地训练后返回的模型权重。每一轮通信开始时所有客户端都从同一个全局模型出发在本地跑若干轮 SGD把更新后的权重返回服务器服务器按样本量加权平均得到新的全局模型。为什么要强调按样本量加权因为不同参与方的数据量往往差很多。如果简单平均样本量大的机构贡献被稀释聚合出来的模型会偏向数据少的机构收敛变慢甚至偏移。另外两个容易被忽略的细节一是客户端必须从同一个初始模型出发不能各自随机初始化二是本地训练时学习率、batch size、本地 epoch 要记录清楚否则聚合出来的模型含义不明确复现时对不上。2.3 用两个客户端手算一遍 FedAvg用一个极端例子说明加权的重要性。假设客户端 1 有 80 个样本客户端 2 只有 20 个样本全局模型是一个单参数模型初始 w0。本地训练后客户端 1 返回 w10.4客户端 2 返回 w2-0.2。按 FedAvg 加权平均import numpy as np n1, n2 80, 20 w1, w2 0.4, -0.2 # 按样本量加权 fedavg_w (n1 * w1 n2 * w2) / (n1 n2) # 简单平均 naive_w (w1 w2) / 2 print(fFedAvg 加权结果: {fedavg_w:.3f}) print(f简单平均结果: {naive_w:.3f})这段代码把聚合逻辑压缩到最小方便理解“加权”二字的分量。运行结果是 0.28 和 0.10两者差了近 3 倍。如果客户端 2 的本地数据质量本身有问题简单平均会把噪声放大。提示本地模拟的第一步不是写训练代码而是先实现聚合公式并用这个小例子验证正确性。聚合错了后面所有客户端训练都是白跑。3. 用 Python 从零实现本地模拟数据切分、客户端训练与服务器聚合3.1 环境准备先解决 Python 解释器和 torch 版本匹配开始写代码前先确认环境。本地模拟只需要 numpy 和 PyTorchCPU 版就够不需要安装任何联邦学习框架。很多新手在“vscode python环境配置”这里翻车最常见的问题是命令行里 pip 装好了 torch但 IDE 里跑代码仍然报 ModuleNotFoundError原因就是 IDE 选了另一个解释器。我一般这样准备环境python -m venv .venv source .venv/bin/activate # Windows 下是 .venv\Scripts\activate pip install numpy torch --index-url https://download.pytorch.org/whl/cpu这里用 venv 而不是 conda是因为联邦模拟的依赖很少虚拟环境越轻越好。torch 指定 CPU 版可以避免下载几百 MB 的 CUDA 依赖本地模拟用不到 GPU。装完后在 IDE 里把解释器指向 .venv 下的 python然后跑一句import torch; print(torch.__version__)验证。3.2 构造模拟数据集两类高斯分布二分类本地模拟的目的是验证流程而非刷点所以我倾向用合成数据而不是 MNIST。合成数据的好处是训练快、维度低模型不收敛时一眼能看出问题。这里生成两类二维高斯分布样本每类 500 条import numpy as np def generate_data(n_samples_per_class500, seed42): rng np.random.default_rng(seed) class0 rng.multivariate_normal( mean[2.0, 2.0], cov[[1.0, 0.2], [0.2, 1.0]], sizen_samples_per_class) class1 rng.multivariate_normal( mean[6.0, 6.0], cov[[1.0, 0.2], [0.2, 1.0]], sizen_samples_per_class) X np.vstack([class0, class1]).astype(np.float32) y np.hstack([ np.zeros(n_samples_per_class, dtypenp.int64), np.ones(n_samples_per_class, dtypenp.int64), ]) idx rng.permutation(len(y)) return X[idx], y[idx] X, y generate_data() print(X.shape, y.shape, y.mean())rng np.random.default_rng(seed)是 NumPy 1.17 之后推荐的随机数风格比旧的np.random.seed更安全。两个类别的中心点距离足够远线性模型也能学到 90% 以上的准确率这样后续加 Non-IID、调参时效果差异不会被随机噪声淹没。y.mean()接近 0.5说明类别均衡。3.3 IID 数据切分按样本 ID 模拟横向参与方横向联邦的数据划分是“按行切分”每个客户端拿到的都是完整特征列的子集。IID 场景下把所有样本随机打乱再均分给客户端def split_iid(X, y, num_clients5, seed42): rng np.random.default_rng(seed) idx rng.permutation(len(y)) client_data [] for part in np.array_split(idx, num_clients): client_data.append((X[part], y[part])) return client_data client_data split_iid(X, y, num_clients5) for cid, (cx, cy) in enumerate(client_data): print(fclient {cid}: {len(cy)} samples, label ratio{cy.mean():.2f})np.array_split会在样本数不能被整除时自动分配余数5 个客户端各拿到 200 条。每个客户端的标签比例都在 0.5 附近这就是 IID各参与方数据分布一致。本地模拟中这个函数决定了“参与方之间到底有多像”后面 Non-IID 实验只需要替换这个函数。3.4 客户端本地训练从同一全局模型出发深拷贝后训练定义一个小 MLP然后实现客户端训练函数。这里最关键的是copy.deepcopy因为 PyTorch 的state_dict()返回的是参数引用直接赋值会导致全局模型被本地训练污染import copy import torch import torch.nn as nn import torch.optim as optim class MLP(nn.Module): def __init__(self, input_dim2, hidden_dim16, output_dim2): super().__init__() self.net nn.Sequential( nn.Linear(input_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, output_dim), ) def forward(self, x): return self.net(x) def client_train(global_model, client_data, local_epochs2, lr0.01, batch_size32): local_model copy.deepcopy(global_model) optimizer optim.SGD(local_model.parameters(), lrlr) loss_fn nn.CrossEntropyLoss() X, y client_data dataset torch.utils.data.TensorDataset(torch.from_numpy(X), torch.from_numpy(y)) loader torch.utils.data.DataLoader(dataset, batch_sizebatch_size, shuffleTrue) local_model.train() for _ in range(local_epochs): for xb, yb in loader: optimizer.zero_grad() logits local_model(xb) loss loss_fn(logits, yb) loss.backward() optimizer.step() return local_model.state_dict(), len(y)copy.deepcopy(global_model)得到一份独立的模型副本本地训练只修改副本返回的state_dict是训练后的参数快照。返回len(y)是为了给服务器端加权聚合提供样本量。这个函数的设计有一个边界要注意DataLoader默认会开多线程加载数据合成数据量很小这里不需要调num_workersWindows 下设置不当反而容易报错。3.5 服务器端 FedAvg 聚合按样本量加权平均服务器拿到所有客户端的参数和样本量后执行 FedAvgdef fed_avg(global_model, client_states, client_sizes): avg_state {} total_size sum(client_sizes) for key in global_model.state_dict().keys(): avg_state[key] sum( state[key] * size / total_size for state, size in zip(client_states, client_sizes) ) global_model.load_state_dict(avg_state) return global_modelstate[key] * size / total_size这一步在 PyTorch 里是张量与标量相乘得到加权后的参数。zip(client_states, client_sizes)保证每个客户端参数与自己的样本量一一对应。聚合完成后用load_state_dict把加权平均结果写回全局模型。注意有些实现会把“返回模型参数”写成“返回梯度更新量”即global_w - local_w。两者在数学上等价但聚合公式不同。本地模拟里建议统一用“返回训练后的完整参数”这种写法语义更清晰排查问题更方便。3.6 主循环通信轮次与指标记录主循环模拟多轮通信。每轮所有客户端从当前全局模型出发本地训练服务器聚合然后评估全局模型效果def evaluate(model, X, y): model.eval() with torch.no_grad(): logits model(torch.from_numpy(X)) pred logits.argmax(dim1).numpy() return (pred y).mean() num_clients 5 num_rounds 15 local_epochs 2 lr 0.01 batch_size 32 global_model MLP() for rnd in range(num_rounds): states, sizes [], [] for cid in range(num_clients): state, size client_train( global_model, client_data[cid], local_epochslocal_epochs, lrlr, batch_sizebatch_size, ) states.append(state) sizes.append(size) fed_avg(global_model, states, sizes) acc evaluate(global_model, X, y) print(fround {rnd 1:2d} | acc{acc:.4f})num_rounds15是通信轮数每轮代表一次完整的“下发-训练-聚合-回传”。local_epochs2是每个客户端在本地数据上完整遍历的次数。这里有个值得注意的现象即使每轮只有 2 个本地 epoch15 轮过后模型准确率也能逼近 95%这是 FedAvg 收敛性的直观体现。如果准确率不升反降优先检查聚合是否按样本量加权、客户端是否从同一全局模型出发。4. 设计可信的模拟实验从 IID 到 Non-IID、三个必调参数4.1 实验基线全局集中训练作为上界本地模拟最容易犯的错误是“只有联邦训练没有基线”。没有基线你没法回答“联邦训练到底损失了多少效果”。常见做法是把同一份数据集中到一块用相同模型训练作为效果上界。def train_centralized(X, y, epochs30, lr0.01, batch_size64): model MLP() optimizer optim.SGD(model.parameters(), lrlr) loss_fn nn.CrossEntropyLoss() dataset torch.utils.data.TensorDataset(torch.from_numpy(X), torch.from_numpy(y)) loader torch.utils.data.DataLoader(dataset, batch_sizebatch_size, shuffleTrue) model.train() for _ in range(epochs): for xb, yb in loader: optimizer.zero_grad() loss loss_fn(model(xb), yb) loss.backward() optimizer.step() return model center_model train_centralized(X, y) center_acc evaluate(center_model, X, y) print(fcentralized acc{center_acc:.4f})集中训练跑 30 个 epoch联邦训练跑 15 轮 x 2 个本地 epoch两者在“总计算量”上接近15x230。这样对比才是公平的如果集中训练 30 epoch 得到 98% 准确率联邦训练 15 轮得到 95%你可以说“联邦只损失了 3 个点”。如果差距超过 10 个点说明聚合实现有问题而不是联邦学习本身不好。4.2 三个必调参数本地 epoch、客户端数量、参与比例参数调优是本地模拟的重头戏。根据我的经验最先要调三个参数参数取值范围影响方向本地 epoch1 ~ 5越大每轮通信效率越高但过大导致客户端漂移客户端数量5 ~ 20越多数据划分越碎单客户端过拟合风险越高参与比例0.5 ~ 1.0模拟客户端掉线或通信资源受限的场景本地 epoch 是最敏感的参数。设成 1 时每轮通信只做一次本地梯度更新行为接近并行 SGD设成 5 或更大每个客户端在自己的数据上反复拟合聚合出来的模型容易震荡。客户端数量影响的是数据切分的粒度5 个客户端每个有 200 条数据20 个客户端每个只有 50 条在数据量少的情况下某些客户端本地训练会过拟合。参与比例则是模拟真实系统的关键每轮只让 60% 的客户端参与聚合看模型是否还能收敛。这个参数在本地模拟里经常被忽略但真实联邦系统中客户端掉线是常态。4.3 用狄利克雷分布生成 Non-IID 数据真实联邦场景中参与方数据分布几乎不可能是 IID。最常用的 Non-IID 模拟方法是狄利克雷分布采样控制每个客户端上每个类别的样本比例def split_non_iid(X, y, num_clients5, alpha0.5, seed42): rng np.random.default_rng(seed) classes np.unique(y) client_ids [[] for _ in range(num_clients)] for cls in classes: cls_idx np.where(y cls)[0] proportions rng.dirichlet([alpha] * num_clients) proportions proportions / proportions.sum() sizes (proportions * len(cls_idx)).astype(int) diff len(cls_idx) - sizes.sum() sizes[-1] diff rng.shuffle(cls_idx) start 0 for cid, size in enumerate(sizes): client_ids[cid].extend(cls_idx[start:start size]) start size result [] for cid in range(num_clients): ids np.array(client_ids[cid]) rng.shuffle(ids) result.append((X[ids], y[ids])) return resultalpha控制 Non-IID 的程度数值越小每个客户端的数据越偏向某些类别alpha0.1时很可能出现某个客户端 80% 样本都是同一类别。rng.dirichlet([alpha] * num_clients)生成的是多项分布的概率向量再映射到样本数上。最后sizes[-1] diff是为了处理整数截断误差保证每个类别的样本全部分配完毕。跑对比实验时固定其他参数不变只改alpha0.1、0.5、10.0。alpha10.0时各客户端分布接近 IIDalpha0.1时效果会明显下滑。如果你发现 Non-IID 场景下联邦训练准确率比集中训练低很多先别急着换算法试试调大本地 epoch 或者增加每轮参与比例往往能缓解。5. 本地模拟避坑5 个让结果翻车的细节5.1 本地 epoch 过大导致模型漂移loss 不降反升现象前几轮准确率正常上涨第 5 轮之后全局模型准确率突然掉 20 个百分点之后反复震荡。原因本地 epoch 设成 10 甚至 20每个客户端在自己的小数据集上反复拟合模型参数偏离全局最优越来越远FedAvg 加权平均把多个“偏了”的模型混在一起得到一个更差的模型。这在数据量小的客户端上尤其明显。解决把本地 epoch 降到 1~3观察收敛曲线。如果 1 个 epoch 收敛太慢优先增加通信轮数而不是加大本地 epoch。另一个缓解手段是降低本地学习率比如从 0.01 降到 0.005让客户端参数漂移幅度变小。5.2 随机种子没固定两次运行结果对不上现象同一套参数昨天跑 95% 准确率今天跑 91%怀疑代码被改坏了。原因数据生成、数据切分、模型初始化、DataLoader 的 shuffle 全都依赖随机数。任何一处随机种子不同整个实验轨迹都会变。这不是模型实现的问题是实验可复现性的问题。解决固定三个位置的种子。数据生成和切分用np.random.default_rng(seed)模型初始化在MLP()前调用torch.manual_seed(seed)DataLoader 的 shuffle 在 PyTorch 内部也受全局种子影响。写实验时把种子作为参数传进所有函数不要分散在代码各处。5.3 聚合时没有按样本量加权小客户端被放大现象对比实验结果异常联邦训练准确率比集中训练还高或者某些标签预测特别差。原因聚合代码写成了sum(state for state in client_states) / num_clients这是简单平均。数据量小的客户端与大客户端平权模型被小客户端带偏。手算例子中 0.28 和 0.10 的差异已经说明问题。解决检查聚合函数是否接收client_sizes参数。经验是打印每个客户端的样本量确认加权系数与样本量成正比。更稳妥的做法是写一个单元测试两个客户端一个返回全 1 参数一个返回全 0 参数样本量 3:1聚合结果应为 0.75用代码验证而不是肉眼检查。5.4 浅拷贝模型导致全局参数被客户端训练污染现象第一轮通信后全局模型准确率异常高后续轮次反而下降结果完全不符合 FedAvg 行为。原因直接local_model global_model没有deepcopy。客户端训练时修改了local_model的参数因为这是同一份对象引用全局模型也被改了。第一轮看起来“假收敛”后续训练没有收敛基础。解决客户端训练函数第一行写copy.deepcopy(global_model)训练和返回只在副本上进行。排查时可以在聚合前打印全局模型参数哈希值确认每一轮开始时全局模型没有被上一轮的本地训练污染。这也是本地模拟最容易踩的 Python 引用陷阱。5.5 Non-IID 切分后客户端缺类loss 变成 nan现象使用狄利克雷分布切分后某个客户端训练时 loss 输出 nan全局模型准确率崩到 50%。原因alpha很小的时候某个客户端可能分不到某个类别的任何样本。二分类任务中客户端全拿到类别 0本地模型学到“永远输出类别 0”的偏置反向传播时梯度异常最终上报的参数把全局模型带崩。解决切分时增加最小样本约束。最简单的方法是在split_non_iid结束后检查每个客户端np.unique(y)是否覆盖全部类别如果缺类从其他客户端匀几个样本过来。更深层的思路是调整alpha下限比如 0.5而不是极端到 0.05。联邦学习真实场景中参与方数据就是可能不完整本地模拟里先跑通完整标签场景再逐步收紧条件。6. 从模拟到真实联邦验证协议一致性、做一次多进程扩容6.1 一致性验证联邦训练与集中训练模型的效果差距本地模拟跑通后第一个要回答的问题是“联邦训练比集中训练差多少”。把两个模型在同一个测试集上做对比差异有两类一类是随机性造成的抖动一类是聚合机制造成的系统性偏差。建议固定种子跑 3 次取均值如果联邦训练与集中训练准确率稳定差在 2% 以内说明算法实现正确。6.2 用多进程模拟通信延迟提前暴露协议问题本地模拟的下一步是把 for 循环换成ProcessPoolExecutor让每个客户端在独立进程中训练。这一步不是性能优化而是模拟真实系统的通信边界提前暴露参数无法序列化、客户端状态无法传递这类问题from concurrent.futures import ProcessPoolExecutor def train_one_client(args): global_state, X, y, local_epochs, lr, batch_size args model MLP() model.load_state_dict(global_state) optimizer optim.SGD(model.parameters(), lrlr) loss_fn nn.CrossEntropyLoss() dataset torch.utils.data.TensorDataset(torch.from_numpy(X), torch.from_numpy(y)) loader torch.utils.data.DataLoader(dataset, batch_sizebatch_size, shuffleTrue) model.train() for _ in range(local_epochs): for xb, yb in loader: optimizer.zero_grad() loss loss_fn(model(xb), yb) loss.backward() optimizer.step() return model.state_dict(), len(y) with ProcessPoolExecutor(max_workers5) as pool: futures [ pool.submit(train_one_client, (global_model.state_dict(), cx, cy, 2, 0.01, 32)) for cx, cy in client_data ] results [f.result() for f in futures]多进程版本中每个子进程必须load_state_dict(global_state)因为子进程无法直接读取主进程的模型对象。如果你发现state_dict无法 pickle说明后续接真实框架时也要解决模型序列化问题。我个人的习惯是本地模拟阶段就把多进程版本跑通一次因为从单进程改成多进程是通往真实联邦系统最近的一步。模拟完记得对比多进程与单进程的聚合结果两者必须完全一致不一致说明状态传递有遗漏。这个习惯帮我提前暴露了很多协议细节问题希望帮到你。本文还有配套的精品资源点击获取
返回列表