跳转至

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 通过上下文管理器优雅地处理了这个切换。