Hugging Face 訓練效率:使用 Flash Attention 2 的打包

Hugging Face 已推出具備邊界感知的指令微調樣本打包,現在已相容於 Flash Attention 2。此更新允許在不使用 padding 的情況下串接序列,提供最高 2 倍的訓練吞吐量提升與 20% 的峰值記憶體使用量減少,且不影響收斂品質。

提升訓練吞吐量與記憶體效率

在不使用 padding 的情況下打包樣本可顯著降低與無關 padding 代幣相關的計算開銷。改善幅度取決於訓練資料集中序列長度的變異程度:

  • 高變異資料集(例如 FLAN): 在 FLAN 資料集上訓練時,包含 Llama 2-7B、Mistral-7B 與 Granite-8B-code 的模型顯示出 2 倍的吞吐量提升。峰值記憶體使用量降低了 20%。
  • 低變異資料集(例如 OrcaMath): 在 OrcaMath 資料集上訓練(該資料集樣本較長且變異較低)時,吞吐量提升了 1.4 倍,峰值記憶體減少了 6%。

由於新實作保留了 minibatch 並維持與使用 padding 訓練相同的優化步數,驗證損失保持一致,確保訓練收斂不會退化。

技術實作:邊界感知

先前的打包實作在使用 Flash Attention 2 時常忽略樣本邊界,導致不期望的跨樣本注意力,損害模型品質。Hugging Face 透過在打包過程中保持邊界感知來解決此問題。

此方式透過向 Flash Attention 2 提供 position_ids 並使用 flash_attn_varlen_func 來實現,該函式會為每個 mini-batch 計算累積序列長度(cu_seqlens)。此方法讓模型能將序列串接成單一張量,同時確保注意力僅限於正確的序列邊界。

支援的模型

此解決方案要求模型公開 position_ids。目前支援 14 種模型,包括:

  • Llama 2 and 3
  • Mistral and Mixtral
  • Granite
  • DBRX
  • Falcon
  • Gemma
  • OLMo
  • Phi 1, 2, and 3 (including phi3)
  • Qwen 2 and 2 MoE
  • StableLM
  • StarCoder 2

整合與使用

使用者可以根據所使用的函式庫,透過兩條主要路徑實作使用 Flash Attention 2 的打包:

使用 Hugging Face Trainer

若要在 Transformers 函式庫的 Trainer 類別中使用此功能,使用者必須:

  1. 使用 attn_implementation="flash_attention_2" 來實例化模型。
  2. 使用 DataCollatorWithFlattening 這個 collator。

使用 TRL SFTTrainer

對於使用 TRL 函式庫的 SFTTrainer 且搭配 DataCollatorForCompletionOnlyLM 的使用者,需求如下:

  1. 使用 Flash Attention 2 來實例化模型。
  2. 在呼叫 DataCollatorForCompletionOnlyLM 時設定 padding_free=True 旗標。

結論

透過消除 padding 代幣並實作具備邊界感知的注意力,Hugging Face 大幅提升了指令微調的效率。當在樣本長度變異度高的資料集上訓練時,吞吐量與記憶體的提升最為顯著。

Sources