跳转至

pi0_utils.py — PI0 模型的数据变换工具集合

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

文件概述

PI0 模型的数据变换工具集合,包括归一化/反归一化、图像变换、文本 tokenization 和动作处理。这些变换构成了 PI0 模型的输入/输出预处理管道。

核心变换类

Normalize / Unnormalize

class Normalize:
    """将数据归一化到 [-1, 1] 范围"""
    def __init__(self, stats, use_quantiles=False):
        if use_quantiles:
            self.low = stats["q01"]   # PI0.5 使用分位数
            self.high = stats["q99"]
        else:
            self.low = stats["min"]   # PI0 使用 min/max
            self.high = stats["max"]

    def __call__(self, x):
        return 2 * (x - self.low) / (self.high - self.low + 1e-8) - 1

class Unnormalize:
    """将 [-1, 1] 范围的数据反归一化回原始尺度"""
    def __call__(self, x):
        return (x + 1) / 2 * (self.high - self.low + 1e-8) + self.low

ImageTransform

class ImageTransform:
    """图像预处理变换"""
    def __init__(self, resize_imgs_with_padding=(224, 224), enable_image_aug=False):
        self.target_size = resize_imgs_with_padding
        self.enable_aug = enable_image_aug

    def __call__(self, images):
        """
        1. 缩放到目标尺寸(保持宽高比,用 padding 补齐)
        2. 归一化到 [0, 1]
        3. 可选的数据增强(训练时)
        """
        ...

PromptTokenizerTransform

class PromptTokenizerTransform:
    """将文本提示转换为 token IDs"""
    def __init__(self, max_length=200, discrete_state_input=False):
        self.tokenizer = AutoTokenizer.from_pretrained("google/paligemma-3b-pt-224")
        self.max_length = max_length

    def __call__(self, prompt: str):
        tokens = self.tokenizer(
            prompt,
            max_length=self.max_length,
            padding="max_length",
            truncation=True,
            return_tensors="pt"
        )
        return tokens["input_ids"], tokens["attention_mask"]

动作处理

class DeltaActions:
    """将绝对动作转换为增量动作(相对于当前状态)"""
    def __call__(self, actions, current_state):
        return actions - current_state

class AbsoluteActions:
    """将增量动作转换回绝对动作"""
    def __call__(self, delta_actions, current_state):
        return delta_actions + current_state

批量变换

class PadStatesAndActions:
    """将状态和动作填充到固定维度

    不同机器人有不同的状态/动作维度,
    统一填充到 max_state_dim / max_action_dim。
    """
    def __call__(self, state, action):
        state = F.pad(state, (0, self.max_dim - state.shape[-1]))
        action = F.pad(action, (0, self.max_dim - action.shape[-1]))
        return state, action

ALOHA 机器人特殊处理

class AlohaInputs:
    """ALOHA 双臂机器人的输入预处理"""
    def __call__(self, data):
        # 处理双臂的状态拼接
        ...

class AlohaOutputs:
    """ALOHA 双臂机器人的输出后处理"""
    def __call__(self, actions):
        # 将拼接的动作分割给两个手臂
        ...

核心类列表

名称 说明
Normalize 归一化到 [-1, 1]
Unnormalize 反归一化到原始尺度
ImageTransform 图像缩放+归一化
PromptTokenizerTransform 文本 tokenization
DeltaActions 绝对 -> 增量动作
AbsoluteActions 增量 -> 绝对动作
PadStatesAndActions 状态/动作维度填充
AlohaInputs/Outputs ALOHA 机器人特殊处理

与其他模块的关系

  • 被 PI0ForActionPrediction(modeling_pi0_torch.py)使用
  • 被 LiberoPi0Input/Output(policy/libero_policy.py)使用

小结

这个文件是 PI0 模型的"数据管道",负责将各种原始数据(图像、文本、状态、动作)转换为模型可以处理的标准格式。归一化/反归一化确保了数值稳定性。