这篇继续聊联邦学习(二)。上一篇文章我们把联邦学习的基本概念、参与角色和典型架构捋了一遍,今天直接进入训练链路。很多人入坑之后会发现,算法书上的联邦学习写得明明白白,真到自己动手跑起来全是问题:模型不收敛、通信太慢、客户端之间数据分布差距大、服务器聚合之后指标反而变差。这篇文章就把训练链路里的关键节点逐个拆开,讲清楚每个环节的“为什么”,再给一份可以直接复现的 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_{t+1} = 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 原文里C=0.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||^2
mu越大,客户端越不敢偏离全局模型。这个方案实现成本低,效果在 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 参数越小,数据分布越偏。我一般用 alpha=0.5 做比较激烈的 non-IID,用 alpha=100 近似 IID。
import numpy as np import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms def partition_mnist_noniid(num_clients=20, alpha=0.5): train_data = datasets.MNIST(root='./data', train=True, download=True, transform=transforms.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_size=32, shuffle=True)) 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, epochs=3, lr=0.01, device='cpu'): model.train() optimizer = torch.optim.SGD(model.parameters(), lr=lr, momentum=0.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_clients=20, alpha=0.5) for r in range(rounds): sampled_ids = np.random.choice(20, size=8, replace=False) weight_list, sample_counts = [], [] for cid in sampled_ids: model_copy = copy.deepcopy(global_model) client_w = client_train(model_copy, client_loaders[cid], epochs=3) 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(f"round {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,降低参与比例,考虑梯度压缩 |
| 异步更新污染全局模型 | 过时更新未衰减、时间戳不完整 | 加时间戳衰减或回退到同步训练 |
这张表是我在多个项目里反复用到的自查清单。每条对应的调试方法,前面章节都已经展开过,这里不再重复。
以我自己带项目过来的经验,联邦学习真正难的地方往往不是某个算法本身,而是你永远不知道哪一轮里哪个客户端上传的是坏更新。把监控做在聚合前,把容错做在架构里,比调参重要得多。每轮记录参与客户端数量、更新范数分布、全局指标变化,这些日志在模型异常时能救命。联邦学习这条路,算法只占一半,另一半是工程韧性。