一批手机或几家机构各自拥有训练数据,又不方便把原始记录集中到同一台服务器。联邦学习会把当前模型发给参与方,让它们在本地数据上计算更新,再把更新交回服务器聚合。下一轮继续下发新模型,训练就在多轮往返中推进。

我重新读了 Google 2017 年的联邦平均论文和 TensorFlow Federated 教程。论文把移动设备看成典型客户端,本地数据量不同,分布也不同,有些设备随时会掉线。这个场景和机房里稳定、整齐的分布式训练差别很大。

一轮训练要经过四个动作

服务器先选出一批可用客户端并下发模型。客户端用自己的数据训练若干步,随后上传参数变化。服务器汇总这些更新,得到下一轮模型。TensorFlow Federated 的自定义算法教程把广播、本地更新、上传和服务器更新写成一轮的基本组成。

客户端多做几步本地训练,可以减少往返次数。Google 的 FedAvg 实验覆盖五种模型和四个数据集,相比同步随机梯度下降,所需通信轮次减少了 10 到 100 倍。这个结果来自当时的实验设置,实际节省取决于设备带宽、本地算力和数据差异。

数据差异会带来麻烦。某些手机只积累一种输入习惯,某家医院的病例结构也可能偏向特定人群。各客户端独自训练后,更新方向可能互相拉扯。本地步数太多时,模型还会偏向当前设备的数据。项目要用真实的非独立同分布数据做模拟,平均切分公开数据只能完成最早的功能验证。

设备参与也不稳定。手机要有电、联网并满足运行条件,机构节点可能在一轮中断开。服务器只能从当时在线的一部分客户端收集更新。选择规则若长期偏向网络好、算力强的设备,训练结果也会偏向这些用户。

原始数据不上传,更新仍可能泄露

模型更新包含本地数据留下的统计信号。攻击者在某些条件下可能从梯度或参数变化推断训练样本。联邦学习减少了集中收集原始数据的需要,却不会自动给出完整的隐私保证。TensorFlow Federated 为此提供安全聚合与差分隐私聚合方案,让服务器只得到汇总结果,或限制单个用户对更新的影响。

企业准备试验时,先确认参与方能不能稳定执行本地训练,再测每轮上传量和失败重试。随后检查聚合前谁能看到单个更新,密钥怎样管理,退出的客户端怎样处理。模型准确率还要按不同客户端群体分别看,不能只报服务器上的一个总分。

联邦学习适合数据天然分散、原始记录难以集中而参与方又能运行训练的场景。它把集中数据的问题换成了通信、客户端管理和更新保护。团队愿意接住这些工程工作,数据留在本地才会成为可验证的事实。

参考资料