Hugging Face と Flower を用いたフェデレーテッドラーニング

Flower を用いたフェデレーテッドラーニングアーキテクチャ

フェデレーテッドラーニングは、複数の分散クライアントと中央サーバー間でグローバルモデルを学習することを可能にします。生データを一箇所に集約する代わりに、各クライアントは自分のデータ上でローカルにモデルを学習し、モデルパラメータのみをサーバーに送信します。サーバーは事前に定義された戦略を用いてこれらのパラメータを集約し、グローバルモデルを更新します。

提供された実装では、プロセスは以下の手順で進みます。

  1. Local Training: クライアントは自分のデータセットを使用してローカル学習を行います。
  2. Parameter Exchange: クライアントは get_parameters メソッドを介して更新されたパラメータをサーバーに送信します。
  3. Global Aggregation: サーバーは FedAvg(フェデレーテッド・アベレージング)などの戦略を用いて全クライアントのパラメータを集約し、各ラウンドで全クライアントの重みの平均としてグローバル重みを定義します。
  4. Model Distribution: サーバーは set_parameters メソッドを介して更新されたグローバルパラメータをクライアントに返します。

技術的実装の詳細

モデルとデータセット

実装では、ベースモデルとして distilBERT (distilbert-base-uncased) を使用し、Hugging Face の AutoModelForSequenceClassification を介して二値シーケンス分類タスク用にロードしています。対象タスクは IMDB データセット に対する感情分析で、モデルは映画の評価が肯定的か否定的かを判定するように学習されます。

Hugging Face のワークフロー

データ準備と学習には標準的な Hugging Face パイプラインが使用されます。

  • Data Handling: datasets ライブラリを使用して IMDB データセットを取得し、AutoTokenizer でトークン化した後、PyTorch の DataLoader オブジェクトにロードします。
  • Training Loop: AdamW オプティマイザを使用した標準的な PyTorch トレーニングループが実装されています。
  • Evaluation: evaluate ライブラリを使用してテストフェーズ中に精度と損失の指標を計算します。

Flower クライアント (IMDBClient)

Hugging Face のモデルと Flower フレームワークを接続するために、flwr.client.NumPyClient を継承したカスタムクライアントクラスが作成されます。このクラスは 4 つの重要なメソッドを実装しています。

  • get_parameters: モデルパラメータを NumPy 配列として抽出し、サーバーへ送信します。
  • set_parameters: サーバーから受け取ったパラメータでローカルモデルの state_dict を更新します。
  • fit: ローカルの学習関数 (train) を実行し、更新されたパラメータと使用したサンプル数を返します。
  • evaluate: ローカルのテスト関数 (test) を実行し、損失と精度の指標を返します。

サーバー設定と集約

フェデレーテッドプロセスを調整するために、特定の集約戦略を持つ Flower サーバーが初期化されます。実装では fl.server.strategy.FedAvg を使用し、fraction_fit=1.0fraction_evaluate=1.0 に設定しているため、全クライアントが各ラウンドの学習と評価に参加します。

分散された指標を処理するために、weighted_average 関数が実装されています。この関数は、各クライアントが提供したサンプル数に基づいて指標に重み付けを行い、グローバルな精度と損失を計算し、代表的なグローバルパフォーマンス指標を保証します。

フレームワークの互換性

提供された例は PyTorch を使用していますが、同じフェデレーテッドラーニングのワークフローは TensorFlow でも実装可能であることがガイドで述べられています。また、Flower のシミュレーション機能 (flwr['simulation']) を利用すれば、Google Colab のような単一環境内でフェデレーテッド環境をエミュレートし、テストに使用できます。

Sources