Skip to main content
QUICK REVIEW

[论文解读] Prototype Helps Federated Learning: Towards Faster Convergence

Yu Qiao, Seong-Bae Park|arXiv (Cornell University)|Mar 22, 2023
Privacy-Preserving Technologies in Data被引用 7
一句话总结

该论文提出了一种基于原型的联邦学习框架,通过在最终训练轮次中聚合客户端的类别原型来增强模型推理性能,从而取代对分类器头的依赖。通过计算并平均本地模型倒数第二层的原型,该方法在 MNIST 和 Fashion-MNIST 的非独立同分布(non-IID)数据设置下,测试准确率至少提高 1%,且收敛速度优于 FedAvg 和 Local 基线方法。

ABSTRACT

Federated learning (FL) is a distributed machine learning technique in which multiple clients cooperate to train a shared model without exchanging their raw data. However, heterogeneity of data distribution among clients usually leads to poor model inference. In this paper, a prototype-based federated learning framework is proposed, which can achieve better inference performance with only a few changes to the last global iteration of the typical federated learning process. In the last iteration, the server aggregates the prototypes transmitted from distributed clients and then sends them back to local clients for their respective model inferences. Experiments on two baseline datasets show that our proposal can achieve higher accuracy (at least 1%) and relatively efficient communication than two popular baselines under different heterogeneous settings.

研究动机与目标

  • 解决由于客户端之间数据分布非独立同分布(non-IID)导致联邦学习中模型推理性能下降的挑战。
  • 减少联邦学习中分类器层在数据异质性条件下引起的预测偏差。
  • 在不修改标准联邦学习训练流程(仅在最终轮次进行调整)的前提下,提升收敛速度和测试准确率。
  • 引入一种通信高效的方案,利用类别原型进行推理,最大限度减少额外开销。
  • 证明在数据异质性环境下,基于原型的推理优于标准分类器的推理。

提出的方法

  • 客户端通过对其本地数据中同一类别样本的倒数第二层特征进行平均,计算每个类别的原型。
  • 在最终的全局轮次中,客户端将模型参数和类别原型一并上传至服务器进行聚合。
  • 服务器使用公式 (3) 对所有客户端提供的原型按类别进行平均,计算全局原型。
  • 推理时,模型根据输入样本在倒数第二层输出与全局原型之间的 L2 距离最小的类别进行预测,如公式 (4) 所定义。
  • 该方法可无缝集成至标准联邦学习流程:仅在最终轮次修改标准通信与推理过程。
  • 在 MNIST 和 Fashion-MNIST 上,使用包含两层卷积和两层全连接层的 4 层 CNN 模型进行评估。

实验结果

研究问题

  • RQ1在非独立同分布(non-IID)数据分布下,基于原型的推理是否能提升联邦学习中的模型准确率?
  • RQ2在最终轮次聚合类别原型是否相比标准联邦学习方法具有更快的收敛速度?
  • RQ3在不同数据偏斜程度下,所提方法与 FedAvg 和 Local 训练相比,在准确率和通信效率方面表现如何?
  • RQ4原型聚合能否缓解分类器层在联邦学习模型中引入的偏差?
  • RQ5基于原型的推理策略在不同数据集和数据异质性水平下是否均有效?

主要发现

  • 在 α = 0.05 的 MNIST 数据集上,所提方法相比 Local 方法准确率提高 18.9%,相比 FedAvg 提高 2.4%。
  • 在 α = 0.1 的 MNIST 数据集上,所提方法相比 Local 方法准确率提高 31.5%,相比 FedAvg 提高 1.0%。
  • 在 α = 0.05 的 Fashion-MNIST 数据集上,所提方法相比 Local 方法准确率提高 1.4%,相比 FedAvg 提高 14.3%。
  • 在 α = 0.1 的 Fashion-MNIST 数据集上,所提方法相比 Local 方法准确率提高 18.4%,相比 FedAvg 提高 7.4%。
  • 在所有测试的非独立同分布(non-IID)设置下,所提方法的测试准确率始终比两个基线方法至少高出 1%。
  • 该方法在每轮通信的测试准确率收敛速度上表现出相对更快的收敛速率,表明其具有更高的通信效率。

更好的研究,从现在开始

从阅读论文到最终审阅,大幅缩短您的研究时间。

无需绑定信用卡

本解读由 AI 生成,并经人工编辑审核。