fsdp_utils.py — FSDP 分布式训练工具¶
文件路径: verl/utils/fsdp_utils.py
文件概述¶
fsdp_utils.py 是 verl 中最大的工具文件之一(约 900 行),提供了 PyTorch FSDP (Fully Sharded Data Parallel) v1/v2 的全套工具函数,包括模型加载、权重分片/聚合、LoRA 支持、混合精度配置等。
背景知识¶
FSDP 是 PyTorch 的一种分布式训练策略。与传统的 DataParallel(每个 GPU 持有完整模型)不同,FSDP 将模型参数分片到多个 GPU 上:
- 前向传播时,按需收集完整参数
- 计算完后立即释放
- 这样每个 GPU 只需要 1/N 的参数内存
FSDP v1 是原始版本,v2(也叫 FSDPv2 或 torch.distributed._composable.fsdp)是重写版本,API 更简洁。
核心功能详解¶
1. 获取 FSDP 包装策略¶
def get_fsdp_wrap_policy(module, config):
"""根据配置返回 FSDP 的包装策略"""
# 策略决定了哪些子模块会被单独包装成 FSDP 单元
# 通常是 TransformerLayer 级别
包装粒度影响内存效率和通信开销的平衡:粒度越细,内存越省,但通信越频繁。
2. 模型加载与权重管理¶
def load_fsdp_model_to_gpu(model, state_dict):
"""将 state_dict 加载到 FSDP 模型中"""
def offload_fsdp_model_to_cpu(model):
"""将 FSDP 模型参数卸载到 CPU"""
def load_fsdp_model_from_cpu(model):
"""将 CPU 上的参数恢复到 GPU"""
在 RLHF 训练中,Actor、Critic、Reference Model 不会同时需要 GPU 内存。verl 通过 offload/reload 机制,让不活跃的模型暂时驻留在 CPU 上,极大节省显存。
3. LoRA 支持¶
def init_fn_with_lora(model, lora_config):
"""为 FSDP 模型添加 LoRA 适配器"""
def get_fsdp_wrap_policy_with_lora(module):
"""获取兼容 LoRA 的 FSDP 包装策略"""
LoRA (Low-Rank Adaptation) 只训练模型的一小部分参数(低秩矩阵)。这些函数确保 LoRA 层在 FSDP 下正确分片。
4. 混合精度配置¶
核心函数/类列表¶
| 函数/类 | 说明 |
|---|---|
get_fsdp_wrap_policy() |
获取 FSDP 包装策略 |
load_fsdp_model_to_gpu() |
加载模型到 GPU |
offload_fsdp_model_to_cpu() |
模型卸载到 CPU |
load_fsdp_model_from_cpu() |
从 CPU 恢复模型 |
init_fn_with_lora() |
LoRA 初始化 |
get_mixed_precision_policy() |
混合精度策略 |
FSDPModule |
FSDP v2 的类型别名 |
与其他模块的关系¶
- 依赖
device.py进行设备操作 - 被
activation_offload.py用来检测 FSDP 层 - 被
checkpoint/fsdp_checkpoint_manager.py用来保存/恢复模型 - 被 Actor/Critic worker 在初始化模型时调用
小结¶
fsdp_utils.py 是 verl 使用 FSDP 训练策略的核心基础设施,管理了模型的分片、加载、卸载和 LoRA 适配等关键操作。理解它有助于理解 verl 如何在有限 GPU 显存下训练大模型。