[论文解读] Federated Learning via Synthetic Data
本文提出在联邦学习中传输合成数据而不是梯度更新,在大幅降低通信成本的同时,模型退化很小。
Federated learning allows for the training of a model using data on multiple clients without the clients transmitting that raw data. However the standard method is to transmit model parameters (or updates), which for modern neural networks can be on the scale of millions of parameters, inflicting significant computational costs on the clients. We propose a method for federated learning where instead of transmitting a gradient update back to the server, we instead transmit a small amount of synthetic `data'. We describe the procedure and show some experimental results suggesting this procedure has potential, providing more than an order of magnitude reduction in communication costs with minimal model degradation.
研究动机与目标
- 在联邦学习中降低客户端上传通信成本。
- 利用合成数据近似标准梯度更新。
- 开发在FL中生成与传输合成数据的实用流程。
- 通过可以稳定优化的改进来增强该方法,并在标准基准数据集上进行评估。
提出的方法
- 在客户端生成合成数据,使其引发的更新与真实局部更新相似。
- 使用基于蒸馏的目标使合成引发的更新与实际更新对齐(参数空间中的平方误差,或函数空间中的KL散度)。
- 计算 updateFromSynthetic 函数,将合成数据解码为服务器更新 g。
- 在全体客户端中聚合更新 g,以形成服务器端更新。
- 使用基于反向传播的程序来优化合成数据,利用海森矩阵向量乘积以提高效率。
- 通过在客户端数据上评估交叉熵来跟踪表现最佳的合成数据集,以稳定优化。
实验结果
研究问题
- RQ1由客户端传输的合成数据是否可以在大幅降低通信成本的同时,近似联邦学习中的标准梯度更新?
- RQ2应如何生成和优化合成数据以最好地模仿真实的FL更新,以及该方法对超参数的鲁棒性如何?
- RQ3在这种合成数据FL框架中,通信、计算和近似质量之间的权衡是什么?
- RQ4是否可以扩展该方法,使其通过合成数据实现服务器端向客户端的传输且成本不高?
主要发现
- 合成数据FL只需要极小比例的上传成本(例如,与传输MNIST CNN的梯度更新相比,浮点数占比为2.4%)。
- 在IID和非IID数据分区上,合成数据方法的模型性能可与标准全梯度传输FL相当。
- 该方法对蒸馏学习率在合理范围内(0.03 至 0.3)表现出鲁棒性。
- 增加合成数据量或计算量可以提高近似度,但收益递减,表明保持平衡的混合规模是最优的。
- 将该方法扩展到服务器向客户端的传输是可行但更具挑战性,能在显著降低下载成本的同时带来一定的准确率权衡(约损失1.5%,但下载成本下降超过90%)。
- 该技术可以通过在客户端数据上用交叉熵跟踪表现最佳的合成数据来实现稳定。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。