fsdp_workers.py — RobActorRolloutRefWorker 扩展了通用的 `ActorRolloutRef...¶
文件路径:
verl/experimental/vla/fsdp_workers.py模块路径:verl.experimental.vla.fsdp_workers
文件概述¶
RobActorRolloutRefWorker 扩展了通用的 ActorRolloutRefWorker,为 VLA 模型添加了推理/训练模式切换(FSDP unshard/reshard)、VLA 模型注册和自定义的 generate_sequences 方法。
核心类:RobActorRolloutRefWorker¶
模式切换机制¶
VLA 模型在训练和推理时需要不同的 FSDP 分片状态:
class RobActorRolloutRefWorker(ActorRolloutRefWorker):
@contextmanager
def rollout_mode(self):
"""切换到推理模式
FSDP 需要 unshard(聚合参数)才能做推理,
因为推理需要完整的模型参数。
"""
with FSDP.summon_full_params(self.model, writeback=False):
self.model.eval()
yield
# 退出后自动 reshard(重新分片)
@contextmanager
def trainer_mode(self):
"""切换到训练模式
训练时使用分片参数,节省显存。
"""
self.model.train()
yield
VLA 模型初始化¶
def init_model(self):
"""初始化 VLA 模型
需要先注册自定义模型到 HuggingFace,然后加载。
"""
# 注册 OpenVLA 和 PI0 模型
from verl.experimental.vla.models import register_vla_models
register_vla_models()
# 使用 HuggingFace 标准接口加载模型
model = AutoModelForVision2Seq.from_pretrained(
config.model_path,
trust_remote_code=True,
)
# 包装为 FSDP 模型
self.model = FSDP(model, ...)
推理方法¶
def generate_sequences(self, data: DataProto):
"""在推理模式下生成动作序列
1. 切换到 rollout_mode(unshard FSDP)
2. 调用 NaiveRolloutRob 生成动作
3. 退出 rollout_mode(reshard FSDP)
"""
with self.rollout_mode():
output = self.rollout.generate_sequences(data)
return output
FSDP 模式切换示意¶
训练模式(分片):
GPU 0: [参数分片 0] GPU 1: [参数分片 1] GPU 2: [参数分片 2]
↓ rollout_mode() (unshard)
推理模式(完整):
GPU 0: [完整参数] GPU 1: [完整参数] GPU 2: [完整参数]
↓ 退出 context (reshard)
训练模式(分片):
GPU 0: [参数分片 0] GPU 1: [参数分片 1] GPU 2: [参数分片 2]
核心类/函数列表¶
| 名称 | 类型 | 说明 |
|---|---|---|
RobActorRolloutRefWorker |
类 | VLA 专用 Actor Worker |
rollout_mode |
上下文管理器 | 切换到推理模式(unshard) |
trainer_mode |
上下文管理器 | 切换到训练模式 |
init_model |
方法 | 初始化并注册 VLA 模型 |
generate_sequences |
方法 | 推理模式下生成动作 |
与其他模块的关系¶
- 继承自
ActorRolloutRefWorker - 使用
register_vla_models()(models/register_vla_models.py)注册模型 - 使用
NaiveRolloutRob(naive_rollout_rob.py)做推理 - 使用
RobDataParallelPPOActor(dp_rob.py)做训练 - 被
RobRayPPOTrainer创建和管理
小结¶
RobActorRolloutRefWorker 解决了 FSDP 训练中推理/训练切换的问题。在分布式训练中,模型参数被分片到多个 GPU 上以节省显存,但推理时需要完整参数。这个 Worker 通过上下文管理器优雅地处理了这个切换。