使用 Hugging Face 與 Flower 的聯邦學習
Hugging Face 提供了一份技術指南,說明如何整合 Flower 框架以啟用針對 Transformer 模型的聯邦學習(FL)。此整合允許在多個客戶端上微調預訓練模型,而無需共享原始資料,提升隱私與資料安全性。
使用 Flower 的聯邦學習架構
聯邦學習使得在多個去中心化的客戶端與中心伺服器之間訓練全域模型成為可能。與其在單一位置聚合原始資料,每個客戶端會在本地使用自己的資料訓練模型,並僅將模型參數傳回伺服器。伺服器再使用預先定義的策略聚合這些參數,以更新全域模型。
在提供的實作中,流程遵循以下步驟:
- 本地訓練:客戶端使用自己的資料集執行本地訓練。
- 參數交換:客戶端透過
get_parameters方法將更新後的參數傳送至伺服器。 - 全域聚合:伺服器使用如
FedAvg(聯邦平均)等策略聚合所有客戶端的參數,將全域權重定義為每輪所有客戶端權重的平均值。 - 模型分發:伺服器透過
set_parameters方法將更新後的全域參數回傳給客戶端。
技術實作細節
模型與資料集
此實作使用 distilBERT(distilbert-base-uncased)作為基礎模型,透過 Hugging Face 的 AutoModelForSequenceClassification 載入,用於二元序列分類任務。目標任務是對 IMDB 資料集 進行情感分析,模型被訓練以判斷電影評分是正面還是負面。
Hugging Face 工作流程
標準的 Hugging Face 流程被用於資料準備與訓練:
- 資料處理:使用
datasets套件取得 IMDB 資料集,接著使用AutoTokenizer進行斷詞,並載入至 PyTorch 的DataLoader物件。 - 訓練迴圈:使用
AdamW優化器實作標準的 PyTorch 訓練迴圈。 - 評估:在測試階段使用
evaluate套件計算準確率與損失指標。
Flower 客戶端 (IMDBClient)
為了將 Hugging Face 模型與 Flower 框架結合,透過繼承 flwr.client.NumPyClient 建立自訂客戶端類別。此類別實作了四個關鍵方法:
get_parameters:將模型參數提取為 NumPy 陣列,以傳送至伺服器。set_parameters:使用從伺服器接收的參數更新本地模型的 state dictionary。fit:執行本地訓練函式(train),並回傳更新後的參數與使用的樣本數量。evaluate:執行本地測試函式(test),並回傳損失與準確率指標。
伺服器設定與聚合
為協調聯邦流程,會以特定聚合策略初始化 Flower 伺服器。此實作使用 fl.server.strategy.FedAvg,其設定 fraction_fit=1.0 與 fraction_evaluate=1.0,表示所有客戶端在每一輪訓練與評估皆參與。
為處理分散式指標,實作了 weighted_average 函式。此函式根據每個客戶端貢獻的樣本數量加權其指標,計算全域的準確率與損失,確保得到具代表性的全域效能衡量。
框架相容性
雖然此範例使用 PyTorch,指南指出相同的聯邦學習工作流程亦可使用 TensorFlow 實作。Flower 的模擬功能(flwr['simulation'])亦可用於在單一環境(如 Google Colab)中模擬聯邦環境,以進行測試。