中文

理解与改进异构数据联邦学习中的模型平均

机器学习 2024-06-03 v4 人工智能

摘要

模型平均是联邦学习(FL)中一种被广泛采用的技术,它聚合多个客户端模型以获得全局模型。值得注意的是,即使在客户端模型以非凸目标函数训练且基于异构本地数据集的情况下,FL 中的模型平均仍能产生更优的全局模型。然而,其成功背后的原理仍鲜为人知。为阐明此问题,我们首先可视化 FL 在客户端与全局模型上的损失景观以展示其几何性质。可视化显示客户端模型在一个公共盆地内包含全局模型,且有趣的是,全局模型可能偏离盆地中心同时仍优于客户端模型。为进一步洞察 FL 中的模型平均,我们将全局模型的期望损失分解为与客户端模型相关的五个因素。具体而言,我们的分析揭示早期训练后全局模型损失主要源于 \textit{i)} 客户端模型在客户端数据集与全局数据集非重叠数据上的损失,以及 \textit{ii)} 全局与客户端模型间的最大距离。基于损失景观可视化与损失分解的发现,我们提出在训练后期对全局模型使用迭代滑动平均(IMA)以减小其偏离期望最小值的程度,同时约束客户端探索以限制全局与客户端模型间的最大距离。我们的实验表明,将 IMA 融入现有 FL 方法显著提升了它们在基准数据集多种异构数据设定下的准确率与训练速度。代码见 \url{https://github.com/TailinZhou/FedIMA}。

关键词

引用

@article{arxiv.2305.07845,
  title  = {Understanding and Improving Model Averaging in Federated Learning on Heterogeneous Data},
  author = {Tailin Zhou and Zehong Lin and Jun Zhang and Danny H. K. Tsang},
  journal= {arXiv preprint arXiv:2305.07845},
  year   = {2024}
}

备注

To appear in IEEE Transactions on Mobile Computing. Code is available at https://github.com/TailinZhou/FedIMA