跳转至

processing_prismatic.py — 定义了 OpenVLA 模型的图像处理器 PrismaticImageProcessor 和统一...

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

文件概述

定义了 OpenVLA 模型的图像处理器 PrismaticImageProcessor 和统一处理器 PrismaticProcessor。前者负责图像变换(基于 TIMM),后者将图像处理器和文本 tokenizer 组合成一个统一接口。

核心类

PrismaticImageProcessor

class PrismaticImageProcessor(BaseImageProcessor):
    """基于 TIMM 的图像预处理器

    对输入图像应用模型训练时使用的相同变换(归一化、缩放等)。
    """
    def __init__(self, image_sizes, timm_model_ids):
        # 从 TIMM 模型获取数据变换配置
        self.transforms = []
        for model_id, size in zip(timm_model_ids, image_sizes):
            data_cfg = timm.data.resolve_data_config(model=model_id)
            transform = timm.data.create_transform(**data_cfg)
            self.transforms.append(transform)

    def preprocess(self, images):
        """应用图像变换"""
        if self.use_fused_backbone:
            # 融合骨干:对同一图像分别应用两套变换,然后拼接
            t1 = self.transforms[0](image)  # SigLIP 变换
            t2 = self.transforms[1](image)  # DINOv2 变换
            return torch.cat([t1, t2], dim=0)  # 6通道
        else:
            return self.transforms[0](image)   # 3通道

PrismaticProcessor

class PrismaticProcessor(ProcessorMixin):
    """统一处理器:组合图像处理和文本 tokenization

    使用方式:
        processor = PrismaticProcessor.from_pretrained("model_path")
        inputs = processor(text="...", images=image)
    """
    def __init__(self, image_processor, tokenizer):
        self.image_processor = image_processor
        self.tokenizer = tokenizer

    def __call__(self, text, images=None):
        """处理文本和图像输入"""
        result = {}

        # 处理文本
        text_features = self.tokenizer(
            text, return_tensors="pt", padding=True
        )
        result["input_ids"] = text_features["input_ids"]
        result["attention_mask"] = text_features["attention_mask"]

        # 处理图像
        if images is not None:
            pixel_values = self.image_processor.preprocess(images)
            result["pixel_values"] = pixel_values

        return result

核心类列表

名称 说明
PrismaticImageProcessor TIMM 图像预处理器
PrismaticProcessor 统一的图像+文本处理器

与其他模块的关系

  • 被 naive_rollout_rob.py 的 process_input 使用
  • 被 register_vla_models.py 注册到 HuggingFace 的 AutoProcessor
  • 使用 configuration_prismatic.py 的骨干配置

小结

处理器将原始的图像和文本转换为模型可以接受的张量格式,是模型推理流水线的"入口"。