Skip to main content
QUICK REVIEW

[論文レビュー] Rethinking Client Drift in Federated Learning: A Logit Perspective

Yunlu Yan, Chun-Mei Feng|arXiv (Cornell University)|Aug 20, 2023
Privacy-Preserving Technologies in Data被引用数 6
ひとこと要約

本稿では、クラスプロトタイプ類似度蒸留を用いて局所モデルとグローバルモデルのログ確率を一致させることで、クライアントドリフトを軽減する新しいフェデレーテッドラーニングフレームワーク、FedCSDを提案する。動的マスクを用いて信頼性の低いグローバルソフトラベルを適応的にフィルタリングすることで、悲観的忘却を低減し、非IIDデータ設定において性能を向上させ、CIFAR-100およびFEMNISTで最先端手法を上回り、最大71.53%の精度を達成する。

ABSTRACT

Federated Learning (FL) enables multiple clients to collaboratively learn in a distributed way, allowing for privacy protection. However, the real-world non-IID data will lead to client drift which degrades the performance of FL. Interestingly, we find that the difference in logits between the local and global models increases as the model is continuously updated, thus seriously deteriorating FL performance. This is mainly due to catastrophic forgetting caused by data heterogeneity between clients. To alleviate this problem, we propose a new algorithm, named FedCSD, a Class prototype Similarity Distillation in a federated framework to align the local and global models. FedCSD does not simply transfer global knowledge to local clients, as an undertrained global model cannot provide reliable knowledge, i.e., class similarity information, and its wrong soft labels will mislead the optimization of local models. Concretely, FedCSD introduces a class prototype similarity distillation to align the local logits with the refined global logits that are weighted by the similarity between local logits and the global prototype. To enhance the quality of global logits, FedCSD adopts an adaptive mask to filter out the terrible soft labels of the global models, thereby preventing them to mislead local optimization. Extensive experiments demonstrate the superiority of our method over the state-of-the-art federated learning approaches in various heterogeneous settings. The source code will be released.

研究の動機と目的

  • データの非IID性と非IIDデータ分布が原因で生じるクライアントドリフトを是正すること。
  • 局所モデルとグローバルモデルのログ確率の乖離が性能低下に与える影響を調査すること。
  • クラスプロトタイプ類似度を用いて局所ログ確率を精錬されたグローバルログ確率に一致させることで、モデルの汎化性能を向上させること。
  • 知識蒸留の過程で低品質なグローバルソフトラベルによる誤った誘導を防ぐこと。
  • 通信効率の高い手法を開発し、グローバル知識を維持しながらローカルデータに適応させること。

提案手法

  • FedCSDは、グローバルクラスプロトタイプとの類似度に基づいて重み付けされたグローバルログ確率と一致するように局所ログ確率を整えるクラスプロトタイプ類似度蒸留機構を導入する。
  • 信頼性の低いグローバルソフトラベルを効果的にフィルタリングするための適応的マスクを採用し、初期学習段階で価値あるが誤った予測を保持する。
  • グローバルモデルの精度に応じてフィルタリング率を動的に調整するため、性能が向上するにつれて過剰なフィルタリングを回避する。
  • 局所ログ確率とグローバルプロトタイプとの類似度に基づく重み付き平均を用いてグローバルログ確率を精錬し、知識伝達の質を向上させる。
  • 局所損失関数は交差エントロピーと、精錬されたグローバルログ確率からの乖離をペナルティ化する蒸留損失の組み合わせで構成される。
  • グローバルモデルのプロトタイプ行列は1ラウンドに1回のみ通信され、通信コストは極めて小さく(O(|Y|²))なる。

実験結果

リサーチクエスチョン

  • RQ1非IIDデータ下でフェデレーテッドトレーニングが進行する中、局所モデルとグローバルモデルのログ確率の差はどのように変化するか?
  • RQ2ログ確率のシフトがフェデレーテッドラーニングにおける性能低下にどの程度寄与しているか?
  • RQ3局所ログ確率とグローバルログ確率を一致させることで、特徴分布シフトを効果的に低減し、汎化性能を向上させられるか?
  • RQ4グローバルソフトラベルの品質がフェデレーテッドラーニングにおける知識蒸留に与える影響は何か?そして、どのように改善できるか?
  • RQ5一部の誤ったソフトラベルを保持する適応的マスクは、厳密なフィルタリングと比較して汎化性能を向上させられるか?

主な発見

  • FedCSDは、β=5のラベルスケイプ下でCIFAR-100で71.53%のテスト精度を達成し、FedAvgおよび他の最先端手法を上回る。
  • 適応的マスクは、誤ったソフトラベルをすべて除去する強力なマスクよりも、グローバルモデルの精度が向上するにつれて一貫して性能を向上させる。
  • FEMNISTでは、FedCSDが特徴分布シフトを顕著に低減し、特徴空間におけるクラスクラスタリングと意思決定境界が改善されている。
  • 特にラベルスケイプ設定において、FedAvgと比較して収束が早く、より安定した学習曲線を示す。
  • 可視化結果から、FedCSDは局所モデルにグローバル知識を保持しており、ローカルデータバイアスへの過剰適合を防いでいることが確認された。
  • 通信コストは最小限であり、1ラウンドに|Y|×|Y|のプロトタイプ行列を1回送信するのみで、スケーラビリティに優れる。

より良い研究を、今すぐ始めましょう

論文の読解から最終レビューまで、研究時間を劇的に削減しましょう。

クレジットカード登録不要

このレビューはAIが作成し、人間の編集者が確認しました。