跳转至

modeling_pi0_torch.py — PI0ForActionPrediction 是 PI0/PI0.5 模型的顶层封装

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

文件概述

PI0ForActionPrediction 是 PI0/PI0.5 模型的顶层封装,继承自 HuggingFace 的 PreTrainedModel 并实现了 SupportSACTraining 接口。它整合了输入/输出变换、Flow Matching 推理和 SAC 算法支持。

架构概览

输入: 图像 + 任务描述 + 机器人状态
      ↓
┌─────────────────────────────────────────┐
│  ImageTransform (resize + normalize)     │
│  PromptTokenizerTransform (text -> ids)   │
│  Normalize (state -> [-1,1])             │
│  ↓                                       │
│  PI0Model (PaliGemma + Expert)           │  ← Flow Matching 去噪
│  ↓                                       │
│  Unnormalize ([-1,1] -> 原始尺度)        │  ← 动作反归一化
└─────────────────────────────────────────┘
输出: 连续动作 (B, n_action_steps, action_dim)

核心功能

1. 初始化与变换管道

class PI0ForActionPrediction(PreTrainedModel, SupportSACTraining):
    def __init__(self, config: PI0TorchConfig):
        super().__init__(config)

        # 输入变换
        self.state_normalize_transform = Normalize(config.state_norm_stats)
        self.action_normalize_transform = Normalize(config.action_norm_stats)
        self.image_transform = ImageTransform(resize_imgs_with_padding=(224, 224))
        self.prompt_tokenizer_transform = PromptTokenizerTransform(max_length=200)

        # 输出变换
        self.action_unnormalize_transform = Unnormalize(config.action_norm_stats)

2. Forward(训练用)

def forward(self, data: DataProto, mode: str = "train"):
    """训练前向传播

    Flow Matching 训练目标:预测速度场 v_t
    损失 = MSE(v_predicted, v_target)
    """
    # 准备输入
    images, img_masks, lang_tokens, lang_masks, state = self._prepare_inputs(data)

    # 采样时间步 t ~ Uniform(0, 1)
    timestep = torch.rand(batch_size)

    # 采样噪声
    noise = self.model.sample_noise(action_shape, device)

    # 构造 x_t = (1-t) * action + t * noise
    x_t = (1 - timestep) * action + timestep * noise

    # 预测速度场 v_t
    v_t = self.model(images, img_masks, lang_tokens, lang_masks, state, x_t, timestep)

    # Flow Matching 损失
    target_v = noise - action
    loss = F.mse_loss(v_t, target_v)
    return loss

3. sample_actions(推理用)

def sample_actions(self, data: DataProto):
    """推理时生成动作

    使用 Euler 方法求解 ODE:从噪声 x_1 积分到动作 x_0
    """
    # 准备输入
    images, img_masks, lang_tokens, lang_masks, state = self._prepare_inputs(data)

    # 调用底层模型的采样方法
    raw_actions = self.model.sample_actions(
        images, img_masks, lang_tokens, lang_masks, state
    )

    # 反归一化
    actions = self.action_unnormalize_transform(raw_actions)
    return actions

4. SAC 支持

PI0 模型实现了 SupportSACTraining 接口,支持 SAC 算法训练:

# Critic 网络(Q 值估计)
self.critic_heads = nn.ModuleList([
    MLP(input_dim=2150, hidden_dims=[1024, 512, 256], output_dim=1)
    for _ in range(head_num)  # Double Q: 2 个 Critic
])

# Target 网络(软更新)
self.target_network_heads = nn.ModuleList([...])

5. Flow-SDE 对数概率(SAC 用)

def _get_logprobs(self, state_features, actions):
    """使用 Flow-SDE 公式计算动作对数概率

    这是 SAC 算法所需的 \(\log \pi(a|s)\)。
    Flow Matching 模型不直接输出概率,需要通过 SDE 公式推导。
    """
    # 使用 Hutchinson's trace estimator 近似 log det Jacobian
    ...

核心类/函数列表

名称 类型 说明
PI0ForActionPrediction 类 PI0 顶层模型封装
forward 方法 训练前向传播(Flow Matching 损失)
sample_actions 方法 推理时生成动作
_get_logprobs 方法 Flow-SDE 对数概率(SAC 用)
sac_forward_critic 方法 SAC Critic 前向
sac_forward_actor 方法 SAC Actor 前向
sac_update_target_network 方法 SAC 目标网络软更新

与其他模块的关系

  • 使用 PI0Model(model/modeling_pi0.py)做核心推理
  • 使用 pi0_utils.py 的变换函数
  • 使用 modules/mlp.py 的 MLP 构建 Critic
  • 实现 SupportSACTraining(sac/base.py)接口
  • 被 register_vla_models.py 注册到 HuggingFace

小结

PI0ForActionPrediction 与 OpenVLA 的根本区别是动作预测方式: - OpenVLA:离散化动作为 token,用语言模型直接预测 - PI0:使用 Flow Matching(连续扩散)生成连续动作

PI0 的优势是原生支持连续动作空间,无需离散化带来的精度损失。