跳转至

register_vla_models.py — 将自定义 VLA 模型注册到 HuggingFace 的 Auto 类系统中

文件路径: verl/experimental/vla/models/register_vla_models.py 模块路径: verl.experimental.vla.models.register_vla_models

文件概述

将自定义 VLA 模型注册到 HuggingFace 的 Auto 类系统中。注册后,可以使用 AutoModelForVision2Seq.from_pretrained() 等标准接口加载这些模型。

关键代码

注册 OpenVLA-OFT 模型

def register_openvla_oft() -> None:
    """注册 OpenVLA OFT 模型和处理器"""
    if _REGISTERED_MODELS["openvla_oft"]:
        return  # 幂等性:避免重复注册

    AutoConfig.register("openvla", OpenVLAConfig)
    AutoImageProcessor.register(OpenVLAConfig, PrismaticImageProcessor)
    AutoProcessor.register(OpenVLAConfig, PrismaticProcessor)
    AutoModelForVision2Seq.register(OpenVLAConfig, OpenVLAForActionPrediction)

    _REGISTERED_MODELS["openvla_oft"] = True

注册 PI0 模型

def register_pi0_torch_model() -> None:
    """注册 PI0 模型"""
    if _REGISTERED_MODELS["pi0_torch"]:
        return

    AutoConfig.register("pi0_torch", PI0TorchConfig)
    AutoModelForVision2Seq.register(PI0TorchConfig, PI0ForActionPrediction)

    _REGISTERED_MODELS["pi0_torch"] = True

统一注册入口

def register_vla_models() -> None:
    """注册所有自定义 VLA 模型"""
    register_openvla_oft()
    register_pi0_torch_model()

HuggingFace Auto 类机制

HuggingFace 的 Auto 类(如 AutoModelForVision2Seq)是一个工厂模式。注册后的使用方式:

# 注册(只需要做一次)
register_vla_models()

# 之后可以用标准接口加载
model = AutoModelForVision2Seq.from_pretrained("path/to/openvla-model")
# HF 会自动根据 config.json 中的 model_type 找到对应的类

核心函数列表

名称 说明
register_openvla_oft 注册 OpenVLA-OFT
register_pi0_torch_model 注册 PI0/PI0.5
register_vla_models 统一注册所有模型

与其他模块的关系

  • 被 fsdp_workers.py 的 init_model 调用
  • 依赖 openvla_oft/ 和 pi0_torch/ 子模块的模型定义

小结

这个文件是 VLA 模型与 HuggingFace 生态系统的桥梁,使得自定义模型可以使用标准的 HF 加载接口。