使用 PyTorch FSDP 微调 Llama 2 70B

Hugging Face 已详细阐述了一种使用 PyTorch Fully Sharded Data Parallelism (FSDP) 微调 Llama 2 70B 模型的方法,利用 Hugging Face Transformers、Accelerate 和 TRL 库。通过在设备间分片优化器状态、梯度和参数,此方法使得在多节点、多 GPU 环境中训练巨型模型成为可能。

克服模型加载期间的 CPU RAM 瓶颈

加载 Llama 2 70B 模型通常需要大量的 CPU RAM;如果节点上的每个进程都加载模型,则可能需要大约 2TB 的 CPU RAM(70B 参数 * 4 字节 * 8 GPU)。为了防止内存不足(OOM)错误,Hugging Face 在 transformersaccelerate 中采用了一种特定的初始化策略:

  1. Meta Device 初始化:模型在所有 rank 上使用 meta 设备创建,意味着它是在没有权重的情况下初始化的。
  2. Rank 0 加载:仅在 rank 0 上加载状态字典。
  3. 空参数分配:其他所有 rank 在 meta 设备上使用 torch.empty() 创建空参数。
  4. 状态广播:通过设置 sync_module_states=True,FSDP 在训练开始前将 rank 0 的权重广播到所有其他 rank。

此方法确保每个节点仅有一个进程将预训练模型加载到 CPU RAM 中,从而在设置阶段大幅减少内存占用。

使用分片状态字典进行高效检查点

在 rank 0 上使用 CPU 卸载的 FULL_STATE_DICT 保存完整的中间检查点经常会导致 NCCL 超时错误和显著延迟。为了解决此问题,Hugging Face 建议在 FSDP 配置中使用 SHARDED_STATE_DICT

  • 中间检查点SHARDED_STATE_DICT 按 GPU 分别保存分片,从而实现更快的保存和训练恢复。
  • 最终模型导出:为了获得用于部署的标准模型状态字典,仅在训练结束前调用 trainer.save_model() 之前,将状态字典类型切换为 FULL_STATE_DICT

优化 VRAM 和训练速度

为了降低计算成本并提高训练速度,该实现采用了两种主要技术:梯度检查点和 Flash Attention。

Flash Attention

标准注意力机制在元素级操作(掩码、softmax 和 dropout)过程中由于冗余的高带宽内存(HBM)读写而经常受内存限制。Flash Attention 通过以下方式进行优化:

  • 内核融合:它将中间步骤保存在 SRAM 中,仅将最终结果写回 HBM 一次。
  • 分块:将 NxN softmax/scores 计算分割成块以适应 SRAM 限制,利用在线 softmax 算法。
  • 重新计算:反向传播重新计算所需值,而不是存储前向传播的整个 NxN 矩阵,从而显著降低内存消耗。

Gradient Checkpointing

启用梯度检查点以进一步减少 VRAM 使用,从而在 70B 参数模型的微调过程中允许使用更大的批次或更长的序列长度。

实施和硬件规格

硬件配置

微调使用了以下硬件:

  • 节点:2 个节点(最少需要 1 个)。
  • GPU:每节点 8 个 A100 (80GB) GPU。
  • 互连:NVLink(节点内)和 Elastic Fabric Adapter(节点间)。
  • 系统 RAM:每节点 1TB。
  • CPU:每节点 96 核心。

训练执行

训练过程使用了带有 FULL_SHARD 策略和 TRANSFORMER_BASED_WRAP 自动包装策略的 accelerate launch 命令。通过使用 bf16 启用了混合精度训练。对于配备 8 个 A100 80GB GPU 的单节点设置,推荐使用 bitsandbytespaged_adamw_32bit 优化器来管理内存。

微调在大约 13.5 小时内完成,使用了模型 meta-llama/Llama-2-70b-chat-hf 和数据集 smangrul/code-chat-assistant-v1

Sources