Skip to main content
QUICK REVIEW

[論文レビュー] Re-Weighted Softmax Cross-Entropy to Control Forgetting in Federated Learning

Gwen Legate, Lucas Caccia|arXiv (Cornell University)|Apr 11, 2023
Privacy-Preserving Technologies in Data被引用数 4
ひとこと要約

本稿では、データの非独立同分布性に起因する深刻な忘却を軽減するため、フェデレーテッドラーニングにおけるクライアントレベルでのソフトマックスログィットの再重み付けを提案する。局所学習中に分布外クラスの勾配を抑制することで、WSMはクライアントドリフトを緩和し、特に高いデータ非独立同分布性と低いクライアント参加率の下でグローバルモデルの性能を向上させる。CIFAR-100では最大7.6%、CIFAR-10では最大1.02%の向上が達成された。

ABSTRACT

In Federated Learning, a global model is learned by aggregating model updates computed at a set of independent client nodes, to reduce communication costs multiple gradient steps are performed at each node prior to aggregation. A key challenge in this setting is data heterogeneity across clients resulting in differing local objectives which can lead clients to overly minimize their own local objective, diverging from the global solution. We demonstrate that individual client models experience a catastrophic forgetting with respect to data from other clients and propose an efficient approach that modifies the cross-entropy objective on a per-client basis by re-weighting the softmax logits prior to computing the loss. This approach shields classes outside a client's label set from abrupt representation change and we empirically demonstrate it can alleviate client forgetting and provide consistent improvements to standard federated learning algorithms. Our method is particularly beneficial under the most challenging federated learning settings where data heterogeneity is high and client participation in each round is low.

研究の動機と目的

  • クライアントが局所データに過学習し、他のクライアントのデータの表現を忘れる、フェデレーテッドラーニングにおける局所クライアントの忘却問題に対処すること。
  • データ非独立同分布性とクライアント参加数の制限がクライアントドリフトを悪化させ、グローバルモデルの性能を低下させることを特定すること。
  • データ共有や追加のモデルパラメータを必要とせず、通信効率の高い方法で忘却を軽減する軽量な手法を提案すること。
  • 各クライアントごとのソフトマックス再重み付けによる交差エントロピー目的関数の変更が、クライアント間での一般化性能を向上させ、収束性を高めることを示すこと。
  • 複数のフェデレーテッドラーニングアルゴリズム(FedAvg、SCAFFOLD、FedProx)とデータセット(CIFAR-10、CIFAR-100)を用い、非独立同分布性とクライアント参加率の変動に応じた評価を実施すること。

提案手法

  • 各クライアントのラベルセット外クラスのソフトマックスログィットに対して、クラス固有の重み付けを施すことで、交差エントロピー損失を計算する前の再重み付けを導入する。
  • クライアントのラベルセット外クラスのログィットを低減することで、局所更新における影響を軽減する、変更された交差エントロピー目的関数を適用する。
  • クラス固有の重みを用いた温度スケーリング付きソフトマックスを用い、訓練の安定性を高めるとともに、分布外クラスの勾配更新を低減する。
  • 通信プロトコルや集約手法を変更せずに、FedAvg、SCAFFOLD、FedProxなどの標準的なフェデレーテッド最適化アルゴリズムに重み付きソフトマックス(WSM)目的関数を統合する。
  • 局所最適化中に干渉を最小限に抑えるために、分布内クラスを優先し、分布外クラスの更新を抑制する重み関数を設計する。
  • データ共有やモデル正則化を超える損失関数内の操作にとどめることで、プライバシー保護型フェデレーテッドラーニングと整合性を保ちながら通信効率を確保する。
Figure 1: Illustration of catastrophic forgetting within client rounds. A global model with knowledge of all classes is sent to all clients participating in a given FL round. Local training increases the client model performance on the client’s local distribution but tends to simultaneously decrease
Figure 1: Illustration of catastrophic forgetting within client rounds. A global model with knowledge of all classes is sent to all clients participating in a given FL round. Local training increases the client model performance on the client’s local distribution but tends to simultaneously decrease

実験結果

リサーチクエスチョン

  • RQ1非独立同分布性のあるフェデレーテッドラーニング環境下で、局所クライアントの忘却がグローバルモデル性能にどの程度悪影響を及えるか?
  • RQ2データ共有やモデル正則化を必要とせず、交差エントロピー損失における各クライアントごとのソフトマックスログィットの再重み付けにより、深刻な忘却を軽減できるか?
  • RQ3実世界のフェデレーテッドラーニングにおける主な課題である高いデータ非独立同分布性と低いクライアント参加率下で、提案手法WSMはどの程度の性能を示すか?
  • RQ4WSMはFedAvgに限らず、複数のフェデレーテッドラーニングアルゴリズムに適用した場合、クライアント間での一般化性能を向上させ、収束性を高めるか?
  • RQ5局所学習ステップ数とクライアント参加率が、WSMによる忘却軽減効果に与える影響は何か?

主な発見

  • WSMは分布外クラスの勾配更新を抑制することで、局所クライアントの忘却を軽減し、より安定的かつ一般化可能なクライアントモデルを実現する。
  • 高いデータ非独立同分布性(α=0.1)下でのCIFAR-100では、FedAvg+WSMが27.4%の精度を達成し、FedAvgの19.8%と比較して7.6%の向上を示した。
  • CIFAR-10では、最適な学習率下でFedAvg+WSMが62.4%の精度を達成し、FedAvgの60.8%と比較して1.02%の向上を示した。
  • クライアント参加率が低い(例:1%)場合、FedAvgとFedAvg+WSMの性能差が顕著に現れ、参加クライアント数が増えるにつれてその差は縮小する。
  • 局所イテレーション数を7から21に増加させた場合、FedAvgでは精度の低下が著しく顕著に現れたが、FedAvg+WSMでは著しく安定した性能を示し、忘却制御の優位性が裏付けられた。
  • WSMはFedAvgに限らず、SCAFFOLDやFedProxに対しても性能向上を示し、フェデレーテッド最適化フレームワーク全体への広範な適用可能性を裏付けた。
Figure 2: CIFAR-10 - Distributions of ten randomly selected clients with data partitioned according to a Dirichlet distribution parameterized by $\alpha=0.1$ .
Figure 2: CIFAR-10 - Distributions of ten randomly selected clients with data partitioned according to a Dirichlet distribution parameterized by $\alpha=0.1$ .

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

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

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

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