Hugging Face 训练效率:使用 Flash Attention 2 的打包

Hugging Face 引入了面向边界的打包用于指令微调示例,并已兼容 Flash Attention 2。此更新允许在不使用填充的情况下连接序列,使训练吞吐量提升最高可达 2 倍,峰值内存使用降低 20%,且不影响收敛质量。

提高训练吞吐量和内存效率

在不使用填充的情况下打包示例可以显著降低与无关填充标记相关的计算开销。改进幅度取决于训练数据集中序列长度的方差:

  • 高方差数据集(例如 FLAN): 在 FLAN 数据集上进行训练时,模型包括 Llama 2-7B、Mistral-7B 和 Granite-8B-code 的吞吐量提升了 2 倍。峰值内存使用降低了 20%。
  • 低方差数据集(例如 OrcaMath): 在 OrcaMath 数据集上进行训练,该数据集示例更长且方差更低,吞吐量提升了 1.4 倍,峰值内存降低了 6%。

由于新实现保留了 minibatch 并保持与填充训练相同的优化步数,验证损失保持不变,确保训练收敛不受影响。

技术实现:边界感知

之前的打包实现经常在使用 Flash Attention 2 时忽略示例边界,导致不期望的跨示例注意力,从而损害模型质量。Hugging Face 通过在打包过程中保持边界感知来解决此问题。

实现方式是向 Flash Attention 2 提供 position_ids 并使用 flash_attn_varlen_func,该函数为每个 mini-batch 计算累计序列长度(cu_seqlens)。此方法使模型能够将序列连接成单个张量,同时确保注意力仅限于正确的序列边界。

支持的模型

该方案要求模型公开 position_ids。目前已支持 14 种模型,包括:

  • Llama 2 和 3
  • Mistral 和 Mixtral
  • Granite
  • DBRX
  • Falcon
  • Gemma
  • OLMo
  • Phi 1、2 和 3(包括 phi3)
  • Qwen 2 和 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 标志。

结论

通过消除填充标记并实现边界感知注意力,Hugging Face 大幅提升了指令微调的效率。当在示例长度方差较高的数据集上进行训练时,吞吐量和内存的提升最为显著。

Sources