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 加载接口。