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 類別中使用此功能,使用者必須:
- 使用
attn_implementation="flash_attention_2"來實例化模型。 - 使用
DataCollatorWithFlattening這個 collator。
使用 TRL SFTTrainer
對於使用 TRL 函式庫的 SFTTrainer 且搭配 DataCollatorForCompletionOnlyLM 的使用者,需求如下:
- 使用 Flash Attention 2 來實例化模型。
- 在呼叫
DataCollatorForCompletionOnlyLM時設定padding_free=True旗標。
結論
透過消除 padding 代幣並實作具備邊界感知的注意力,Hugging Face 大幅提升了指令微調的效率。當在樣本長度變異度高的資料集上訓練時,吞吐量與記憶體的提升最為顯著。