联邦学习是一种机器学习技术,其中共享模型在众多分散的设备或服务器上进行训练,这些设备或服务器各自持有本地数据,而原始数据从不被集中。参与者不是将数据上传到中央服务器,而是在本地计算模型更新,并仅发送该更新,该更新与其他参与者的更新聚合,以改进共享的全局模型。
历史与动机
谷歌在2016年的一篇博客文章及随附论文中引入了联邦学习这一术语和实用系统,其动机来自其移动键盘应用Gboard,该应用希望利用用户输入的内容改进下一个词预测,而无需将通常敏感的文本上传到谷歌服务器。该技术解决了深度学习时代对大规模训练数据的需求与日益增长的数据隐私和监管担忧之间的紧张关系,后者在欧盟GDPR等法律中得到了正式化。它提供了一种从边缘设备本身(如手机、医院本地服务器或银行分支机构)生成的数据中获益的方式,而无需承担集中这些数据的法律、后勤或伦理负担。
工作原理
在标准联邦平均算法中,中央服务器将当前全局神经网络模型发送给一部分参与设备。每台设备使用其本地数据通过普通梯度下降方法对模型进行短暂训练,然后将生成的模型更新(而非数据)发送回服务器。服务器对所有参与设备的更新进行平均,以生成改进后的全局模型,并循环重复多轮。由于原始数据从不离开设备,联邦学习在结构上比集中式训练具有隐私优势,但它本身并非隐私无懈可击;通常会在其上叠加差分隐私和安全聚合等技术,以防止模型更新本身泄露有关单个用户的信息,这一担忧与更广泛的AI安全和数据治理工作密切相关。
应用与局限
联邦学习已部署在谷歌的Gboard键盘预测和Android的歌曲识别功能中,并在医疗保健领域得到探索,医院希望在不跨机构共享患者记录的情况下协作训练诊断模型,以及在金融领域用于跨银行欺诈检测。其局限性显著:设备在计算能力、连接性和数据量方面差异很大,这可能导致聚合模型偏向最活跃或连接最好的参与者;通信成本不容忽视,因为每轮都必须传输模型更新;在众多不可靠、低功耗设备上协调训练比在单一受控良好的集群上训练更为复杂。截至2020年代中期,联邦学习仍是一种在数据无法或不应集中时使用的专门技术,而非传统集中式训练大型模型(如大型语言模型)的默认替代方案,后者仍主要在集中式、精选的数据集上训练,这些数据集主要来源于Common Crawl等资源。