ARTICLE DETAIL

资讯详情

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

联邦学习实战:训练链路、FedAvg聚合与non-IID工程实践

联邦学习实战:训练链路、FedAvg聚合与non-IID工程实践 这篇继续聊联邦学习二。上一篇文章我们把联邦学习的基本概念、参与角色和典型架构捋了一遍今天直接进入训练链路。很多人入坑之后会发现算法书上的联邦学习写得明明白白真到自己动手跑起来全是问题模型不收敛、通信太慢、客户端之间数据分布差距大、服务器聚合之后指标反而变差。这篇文章就把训练链路里的关键节点逐个拆开讲清楚每个环节的“为什么”再给一份可以直接复现的 PyTorch FedAvg 代码最后把工程里容易踩的坑过一遍。1. 联邦学习的完整训练链路从一个本地迭代说起大多数教程讲联邦学习都会从“数据不动模型动”这个口号开始但真到动手阶段这个口号帮不上什么忙。我更愿意把联邦学习看成一套分布式训练协议全局模型在服务器训练数据在客户端每一轮训练都是“下发-本地训练-回传-聚合-更新”的循环。理解这个循环的每个动作后面所有的调参、排错才有基础。1.1 一次联邦轮次里的四个动作一个完整的联邦轮次大致分四步。第一步服务器按策略选一批客户端。注意不是所有客户端都参与。现实环境里客户端动辄成千上万全部参与一来通信负担太重二来掉线的客户端会拖垮整轮所以默认是随机采样一部分比如每轮选 10% 到 30%。参与比例是第一个隐藏的重要超参。第二步服务器把当前全局模型下发到被选中的客户端。很多实现里会顺带把训练配置也下发比如本地训练轮数、学习率、批次大小。这里有一个容易忽略的细节下发的必须是模型权重而不是训练好的梯度。客户端拿到的应该是一份可以在本地独立前向反向的完整模型否则后续聚合根本没法做。第三步客户端在本地数据上训练若干轮。这一步和普通深度学习训练几乎一样区别在于数据是私有的、不能离开本地。客户端训练完成后不传原始数据只传模型更新通常是参数差值或新权重。第四步服务器收集所有参与客户端的模型更新做加权聚合更新全局模型然后进入下一轮。聚合策略是整个联邦学习的核心后面单独讲。这四个动作听起来简单但每一环都有不少工程细节会坑人。尤其是第三步和第四步之间客户端返回的更新尺度差异往往比想象中大得多如果直接在服务器端做普通平均很容易被一两个数据量大的客户端带偏。1.2 通信是瓶颈为什么“本地多算一点”更值接着上面的链路说一个反复被验证的经验通信开销经常比计算开销更昂贵。在传统分布式训练里GPU 集群之间的内网带宽很充足节点之间传梯度基本不心疼。联邦学习的客户端则是手机、边缘盒子、机构服务器这类设备网络环境天差地别有的走 WiFi有的走 4G有的干脆是隔几天才能连一次网。如果把一千个客户端每轮都传一次完整模型很多真实场景根本跑不起来。所以 FedAvg 的设计里有个很重要的思想让客户端本地多迭代几轮减少通信轮次。这个思路就像在办公室里与其每个人每改一行字就往主管那儿跑一趟不如各自把手里的活做完一整块再集中汇报一次。本地 epoch 次数和本地批次大小直接影响通信频率epoch 越大通信越少但太大又会让客户端模型过分偏向本地数据聚合效果反而变差。这个平衡没有固定答案我一般从本地 1 轮或 5 轮开始试配合学习率一起调。2. 聚合算法拆解FedAvg 为什么是默认起点在动手写联邦学习代码之前建议先把聚合这件事想明白。目前大部分联邦学习项目的第一版都是 FedAvg不是因为它最聪明而是因为它在简单性和效果之间取得了一个不错的平衡点。2.1 FedAvg 的核心公式和朴素实现FedAvg 的思想很朴素全局模型的新权重等于所有参与客户端权重的加权平均。第 t 轮结束后全局权重更新为W_{t1} sum_k (n_k / n_total) * W_k其中n_k是客户端 k 的本地样本数量n_total是参与客户端样本总量W_k是客户端 k 本地训练后的新权重。注意参与加权的是样本量不是客户端个数。如果 A 客户端有 10000 条数据B 客户端只有 100 条对全局模型的影响自然应该 A 更大否则 B 的更新会被 A 淹没或者反过来 B 的少量噪声也被放大。这个加权设计本质上是在尽量接近“把数据集中到一起训练”的统计效果。对应的朴素实现也很短def fedavg_aggregate(weight_list, sample_counts): weight_list: 每个客户端返回的模型状态字典列表 sample_counts: 每个客户端的样本数列表 total_samples sum(sample_counts) aggregated {k: torch.zeros_like(v) for k, v in weight_list[0].items()} for client_w, count in zip(weight_list, sample_counts): ratio count / total_samples for k in aggregated: aggregated[k] client_w[k] * ratio return aggregated这里初始化用了zeros_like实际操作中我更推荐把第一个客户端的权重直接乘上对应比例再累加避免不必要的深拷贝。这个函数是后续所有聚合实验的地基建议放在单独模块里并加上单元测试。2.2 数据不平衡时的加权聚合陷阱样本量加权听起来天经地义实际上一旦客户端的数据分布差异很大简单加权会出现两个问题。第一个是学习率缩放问题。客户端本地训练时如果直接使用服务器下发时的原始学习率那么数据量大的客户端本地走的方向可能过于激进。因为它的样本多在同样的本地 epoch 下模型权重移动得更远聚合回来的权重更大下一轮全局模型容易被它带偏。常见补救方式是让客户端在本地训练时对学习率做缩放比如按sqrt(1/n_k)或者按目标参与样本量来调整。更稳的做法是先在均匀独立同分布IID模拟数据上跑通再逐渐引入 non-IID这样能清楚看到是分布问题还是聚合问题。第二个是参与客户端的样本量差异极大时的数值稳定性。如果某个客户端样本数只有个位数它的本地训练基本是噪声却还占了一个加权比例。我遇到过极端情况几百个客户端里有一个样本数特别少但本地学习率很高返回的更新范数是别人的几十倍聚合后全局模型直接发散。排查手段是把每轮各客户端的更新范数记下来做成曲线看分布一旦发现离群值优先查这个客户端的样本量和本地 epoch。2.3 几个直接影响收敛的超参这里把我在实验里觉得最有影响的几个超参列一下每个都说清楚为什么。参与比例C每轮参与客户端占总数的比例。C太小全局模型看到的“新鲜数据”太少收敛慢C太大通信压力大且部分慢客户端会拖慢整轮。FedAvg 原文里C0.1就能有不错效果实际项目中建议从 0.1 开始调观察验证集指标和每轮耗时。本地 epochE客户端本地训练轮数。E决定客户端模型在本地数据上“走多远”。E过小则每轮全局更新幅度小通信次数多E过大则客户端模型过分拟合本地分布聚合模型容易震荡。这个参数要结合客户端数据量一起看数据量小的客户端E建议小一点。本地批次大小B影响本地梯度的噪声。B太大会让客户端更新偏保守B太小会让更新噪声大聚合时方差更高。通常沿用单机训练时比较合适的B但需要配合学习率调整。服务端学习率lr_server很多人忽略服务器端其实也可以设置一个学习率对聚合后的更新做缩放。这个参数可以控制全局模型每次更新的步长在客户端本地训练已经用了不小学习率的情况下服务端学习率一般取 1.0 或略小于 1.0比如 0.9。调优时如果发现全局指标震荡优先降低服务端学习率而不是客户端学习率。3. “数据非独立同分布”才是联邦学习的灵魂痛点如果把联邦学习的所有困难排个名非独立同分布non-IID数据一定是第一名。很多人在模拟实验里用随机切分的数据跑出来效果很好一换到真实分布就崩原因就是数据分布变了。3.1 数据分布漂移和本地漂移non-IID 具体指什么简单说每个客户端的数据分布和全局分布不一样。比如手写数字识别任务里客户端 A 可能全是数字 0 到 4客户端 B 全是 5 到 9而全局模型希望学会所有数字。用本地数据训练出来的模型方向自然会偏离全局最优方向这就叫做本地漂移。打个比方几个部门各自对着自己那部分客户需求做产品等把方案汇总到总部时方案之间互相打架。每个部门都在自己那片小天地里“优化”却没有看过全局需求分布。联邦学习里客户端本地训练越久这种偏离越明显这正是 local epoch 不能无限加大的深层原因。服务器端的全局模型也不是一帆风顺。即使每个客户端都认真训练聚合后全局模型依然可能朝某个客户端多的方向偏因为参与客户端的样本分布并不是全局分布的忠实样本。如果某些客户端数据总是多、总被选中全局模型就会慢慢偏向它们。3.2 应对策略FedProx、SCAFFOLD 和更工程化的分组建模学术界针对 non-IID 提出的方法很多我挑几个工程上真正有线索的说说。FedProx 的思路是在客户端本地损失函数上加一个近端项把本地模型拉向全局模型防止客户端跑得太远。实现上就是在损失函数里加上一个 L2 正则loss_total loss_local (mu / 2) * ||w_local - w_global||^2mu越大客户端越不敢偏离全局模型。这个方案实现成本低效果在 non-IID 明显时非常有效。需要注意的是mu太大会让本地训练失去意义更新和没训练差不多需要配合验证集调。SCAFFOLD 更复杂它引入控制变量来估计客户端与全局之间的更新方向差在本地更新时做校正。数学上更优雅但实现难度和存储成本都更高每个客户端要额外维护控制变量。我的建议是小规模实验可以玩真正上线前先想清楚这些控制变量的存储和更新机制是否扛得住。比算法更重要的是工程策略。很多业务场景里客户端其实可以按“领域”或“设备类型”分组比如按手机系统版本、按地区、按业务线分组组内数据分布相对均匀组间差异再通过联邦学习去弥合。这样既降低 non-IID 程度又能在某组模型效果不好时单独回退。我在项目里做过类似设计先把客户端按设备类型分成两组分别跑联邦循环效果比全部混在一起稳定不少。4. 灾难性遗忘联邦场景里容易被忽视的记忆坑灾难性遗忘这个词原先是连续学习里的经典问题模型在学习新任务时把旧任务的参数覆盖掉导致旧任务精度急剧下降。联邦学习里也经常遇到这个问题而且症状更隐蔽。4.1 本地用户的灾难性遗忘先看客户端本地。如果某个用户数据一直在变比如推荐场景里用户兴趣从 A 类内容转向 B 类内容客户端本地模型在持续用新数据训练时会慢慢忘掉旧的兴趣特征。这不是联邦学习特有的但在联邦场景里更难发现因为本地模型很少做完整评估。更常见的情况是“客户端周更新”。比如手机上的输入法模型用户的使用习惯每周都在变如果客户端持续用最近几天的数据做本地训练旧习惯的预测能力会快速下降。应对策略一般是在客户端本地保留一个小型历史数据池每轮训练时混入一部分上一轮的数据或者加知识蒸馏正则让当前模型不要偏离上一轮模型太多。4.2 全局模型也在经历连续学习联邦学习的全局模型其实也在做连续学习每一轮客户端上传的更新都是基于各自本地新数据的“新知识”。当全局模型把这些知识聚合起来时如果各客户端之间的更新方向冲突太大或者某些任务只在特定时间出现全局模型就可能出现灾难性遗忘。一个典型案例是季节性数据。假设做的是供应链销量预测上半年客户端数据全是夏季商品下半年全是冬季商品。全局模型如果一直追着新数据跑到冬天它会把夏季商品的预测能力忘得差不多。等到明年夏天模型又要重新学一遍效率极低。这种全局层面的遗忘处理起来比本地更棘手因为它需要服务器存储旧数据或旧任务信息而联邦学习的意义正是“不收集数据”。所以现实可落地的方案一般是服务端保存上一轮全局模型或模型快照聚合时加一个权重衰减项限制新全局模型偏离旧全局模型太远对更新做方向约束比如在聚合时抑制那些和主流方向差异过大的客户端更新避免少数新任务把全局模型带偏按任务或时间窗口维护多个全局模型副本做任务层级的模型融合新模型只在新增任务上微调。4.3 经验回放与正则化在联邦场景里的落地经验回放在单机连续学习里是拿旧数据样本重新训练联邦场景里不能那么干但可以用“梯度记忆”的思路服务端存一份全局模型的历史梯度或者每个客户端保留自己上一轮的模型快照在本地训练时用正则化项拉近新旧模型。我在一个图像分类联邦项目里用过跨轮正则化客户端本地训练时除了监督损失额外加一个 KL 散度损失约束当前模型对旧样本的输出分布和上一轮模型不要太远。这里的旧样本可以直接从本地缓存里取也可以从服务端下发一个公共代理数据集。效果上在客户端本地数据频繁变化的场景里全局模型在新任务上的精度没有下降多少旧任务精度也不再陡降。这个手段实现起来不复杂但需要警惕一点正则化过强会拖慢新任务的学习速度。具体强度建议以“新增任务指标提升不超过旧任务指标下降”作为基准来调。5. 实操用 PyTorch 跑通一个最小 FedAvg Demo理论聊多了容易飘下面给一个能直接跑的最小 Demo。我会用 MNIST 做例子但故意把数据切成分片模拟 non-IID 场景。整体代码量不大跑熟之后可以替换成自己的数据集和模型。5.1 用 Dirichlet 分布模拟 non-IID 数据划分模拟 non-IID 最常用的方法是用狄利克雷分布给每个客户端分配不同类别的概率。alpha 参数越小数据分布越偏。我一般用 alpha0.5 做比较激烈的 non-IID用 alpha100 近似 IID。import numpy as np import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms def partition_mnist_noniid(num_clients20, alpha0.5): train_data datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransforms.ToTensor()) labels np.array(train_data.targets) client_data {i: [] for i in range(num_clients)} for cls in range(10): idx_cls np.where(labels cls)[0] proportions np.random.dirichlet(alpha[alpha] * num_clients) proportions proportions / proportions.sum() np.random.shuffle(idx_cls) assigned_counts (len(idx_cls) * proportions).astype(int) start 0 for cid, cnt in enumerate(assigned_counts): client_data[cid].extend(idx_cls[start:start cnt].tolist()) start cnt client_loaders [] for cid in range(num_clients): indices client_data[cid] if len(indices) 0: indices [0] subset torch.utils.data.Subset(train_data, indices) client_loaders.append(DataLoader(subset, batch_size32, shuffleTrue)) return client_loaders这段代码有一个可以优化的地方分配样本时直接把剩下的样本全部给最后一个客户端会有样本浪费或倾斜更稳妥的做法是用多项式采样逐样本分配。不过作为 Demo 已经够用了重点是让每个客户端的数据分布明显不同。5.2 客户端训练函数与服务器聚合主循环客户端训练和单机训练几乎一样只需要注意返回的是模型权重而不是 loss。import copy import torch.nn as nn import torch.nn.functional as F class SimpleCNN(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(1, 32, 3, 1) self.conv2 nn.Conv2d(32, 64, 3, 1) self.fc1 nn.Linear(64 * 5 * 5, 128) self.fc2 nn.Linear(128, 10) def forward(self, x): x F.relu(self.conv1(x)) x F.max_pool2d(x, 2) x F.relu(self.conv2(x)) x F.max_pool2d(x, 2) x torch.flatten(x, 1) x F.relu(self.fc1(x)) return self.fc2(x) def client_train(model, loader, epochs3, lr0.01, devicecpu): model.train() optimizer torch.optim.SGD(model.parameters(), lrlr, momentum0.9) for _ in range(epochs): for x, y in loader: x, y x.to(device), y.to(device) optimizer.zero_grad() out model(x) loss F.cross_entropy(out, y) loss.backward() optimizer.step() return model.state_dict() device cuda if torch.cuda.is_available() else cpu global_model SimpleCNN().to(device) rounds 20 client_loaders partition_mnist_noniid(num_clients20, alpha0.5) for r in range(rounds): sampled_ids np.random.choice(20, size8, replaceFalse) weight_list, sample_counts [], [] for cid in sampled_ids: model_copy copy.deepcopy(global_model) client_w client_train(model_copy, client_loaders[cid], epochs3) weight_list.append(client_w) sample_counts.append(len(client_loaders[cid].dataset)) global_w fedavg_aggregate(weight_list, sample_counts) global_model.load_state_dict(global_w) print(fround {r} done, sampled {len(sampled_ids)} clients)这里有个容易被新手忽略的点客户端训练时一定要 deep copy 全局模型否则在客户端上修改的是全局模型本身。我在最初写 Demo 时因为少写了copy.deepcopy导致聚合前后全部张量共享存储参数被反复污染调了好久才发现。5.3 你大概率会观察到的现象和含义跑完这个 Demo你大概率会看到两种情况。第一种是全局 loss 稳步下降但速度比单机训练慢很多。这正常。因为每一轮只有少部分客户端参与而且 non-IID 让更新方向方差更大。想加速可以增大参与比例或者把本地 epoch 调大到 5。第二种是 loss 一开始下降后面在某个水平震荡甚至反弹。这个大概率是学习率过大或聚合没有任何正则。我建议把服务端学习率从 1.0 调到 0.7或者给客户端本地损失加上 FedProx 的近端正则项震荡通常会缓解。更隐蔽的情况是全局模型在训练集上看起来很好但用独立测试集评估发现精度很低。这种通常意味着模型过拟合到了参与客户端的局部分布上需要考虑增加参与客户端的多样性或者减少本地 epoch。6. 工程化避坑清单从 Demo 到真实系统的距离Demo 跑通只是第一步真实系统中的联邦学习难点几乎全在工程。这里把我踩过的一些坑整理成清单按优先级排列。6.1 聚合前的完整性检查客户端掉线、上传损坏、返回格式异常这些在生产环境里不是万一而是日常。不可靠的更新一旦混进聚合轻则一轮训练报废重则全局模型直接发散。我建议在聚合之前至少做三件事检查每个客户端是否返回了模型权重且 key 结构和全局模型一致检查权重的数值是否在合理范围内比如是否有 NaN 或超大值检查更新范数与历史分布的偏差超过一定阈值直接丢弃该客户端。这三步在任何联邦框架里都应该做掉不要等到模型炸了再去翻日志。6.2 通信压缩与异步更新的取舍通信是联邦学习最贵的资源之一。常用的压缩手段有梯度量化、稀疏化、低秩分解。其中梯度稀疏化最容易落地只上传绝对值最大的前 1% 或 5% 的梯度其他置零。但要注意稀疏化会改变更新分布FedAvg 的加权平均需要重新考虑。异步更新是另一个工程方向。同步联邦每一轮要等所有客户端返回异步联邦则允许不同客户端在不同时间点被聚合。好处是系统吞吐量上去了坏处是全局模型可能被过时很久的更新污染需要给每个客户端的更新打上时间戳并做衰减。实际项目中我一般先做同步联邦把通信效率问题用参与比例和本地 epoch 压下去等业务稳定后再考虑异步。6.3 安全聚合和差分隐私的取舍很多业务场景会同时要求“数据不出域”和“模型不泄露隐私”于是安全聚合和差分隐私成了标配。安全聚合保证服务器在聚合过程中无法看到单个客户端的精确更新差分隐私则通过在梯度上添加噪声让攻击者难以反推某个样本是否存在。这两者需要一起用但都要付出代价。安全聚合增加通信轮次和计算量差分隐私的噪声会直接伤害模型效果。我的建议是先跑一个不加隐私机制的基础版本确认模型效果基线再逐步加入噪声观察精度下降幅度如果精度下降超过可接受范围就需要重新评估数据类型和隐私预算分配。没有一个魔法参数能同时满足两边这是个工程权衡题。7. 常见问题速查表最后把联邦学习里高频出现的问题整理成一张速查表方便你排查的时候直接对号入座。7.1 故障排查的优先顺序我在排查联邦学习训练问题时通常按这个顺序来先看每一轮的参与客户端数量和更新范数分布确认没有“坏更新”再看全局验证集指标曲线区分是发散还是不收敛最后才去调超参。顺序调反的话很容易陷入“调了半天 learning rate其实问题出在某个客户端上传了 NaN”的尴尬局面。7.2 高频问题与排查方向现象可能原因优先排查方向全局模型不收敛学习率过大、non-IID 严重、参与客户端太少降服务端学习率增加参与比例调小本地 epoch模型震荡或反弹客户端更新方向冲突、部分客户端数据倾斜严重检查更新范数离群值加 FedProx 近端项训练很快但测试集效果差过拟合到参与客户端的局部分布增加参与客户端多样性减少本地 epoch某一类任务精度突然下降灾难性遗忘新任务淹没了旧任务加跨轮正则化保留模型快照考虑任务分组聚合后权重出现 NaN客户端本地 loss 爆炸、上传格式异常检查客户端数据是否有空样本、loss 是否有除零通信耗时远超计算本地 epoch 太少、参与比例太高增大本地 epoch降低参与比例考虑梯度压缩异步更新污染全局模型过时更新未衰减、时间戳不完整加时间戳衰减或回退到同步训练这张表是我在多个项目里反复用到的自查清单。每条对应的调试方法前面章节都已经展开过这里不再重复。以我自己带项目过来的经验联邦学习真正难的地方往往不是某个算法本身而是你永远不知道哪一轮里哪个客户端上传的是坏更新。把监控做在聚合前把容错做在架构里比调参重要得多。每轮记录参与客户端数量、更新范数分布、全局指标变化这些日志在模型异常时能救命。联邦学习这条路算法只占一半另一半是工程韧性。
返回列表