跳转至

modeling_pi0.py — PI0Model 是 PI0 的核心推理模型

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

文件概述

PI0Model 是 PI0 的核心推理模型,基于 PaliGemma(视觉语言模型)+ Expert(动作专家网络)的双流架构,使用 Flow Matching 进行动作预测。

架构图

┌──────────────────────────────┐
│               actions        │
│               ^              │
│              ┌┴─────┐        │
│  kv cache    │Gemma │        │
│  ┌──────────>│Expert│        │
│  │           │      │        │
│ ┌┴────────┐  │x 10  │        │  ← 10 步去噪
│ │         │  └^──^──┘        │
│ │PaliGemma│   │  │           │
│ │         │   │  robot state │
│ │         │   noise          │
│ └^──^─────┘                  │
│  │  │                        │
│  │  image(s)                 │
│  language tokens             │
└──────────────────────────────┘

核心方法

forward - 训练用前向传播

def forward(self, images, img_masks, lang_tokens, lang_masks,
            state, x_t, timestep):
    """一步去噪预测

    将前缀(图像+文本)和后缀(状态+噪声动作+时间步)
    送入双流 Transformer,预测速度场 v_t。
    """
    # 1. 编码前缀(图像+文本)
    prefix_embs, prefix_pad_masks, prefix_att_masks = \
        self.embed_prefix(images, img_masks, lang_tokens, lang_masks)

    # 2. 编码后缀(状态+动作+时间步)
    suffix_embs, suffix_pad_masks, suffix_att_masks, adarms_cond = \
        self.embed_suffix(state, x_t, timestep)

    # 3. 构造注意力 mask
    att_2d_masks = make_att_2d_masks(pad_masks, att_masks)

    # 4. 双流 Transformer 前向
    (_, suffix_out), _ = self.paligemma_with_expert.forward(
        inputs_embeds=[prefix_embs, suffix_embs], ...
    )

    # 5. 提取动作输出
    v_t = self.action_out_proj(suffix_out[:, -self.n_action_steps:])
    return v_t

sample_actions - 推理用采样

@torch.no_grad()
def sample_actions(self, images, img_masks, lang_tokens, lang_masks, state):
    """完整推理:从噪声生成动作

    使用 Euler 方法求解 ODE:x_{t+dt} = x_t + dt * v_t
    """
    # 1. 编码前缀并缓存 KV
    prefix_embs, ... = self.embed_prefix(images, img_masks, lang_tokens, lang_masks)
    _, past_key_values = self.paligemma_with_expert.forward(
        inputs_embeds=[prefix_embs, None], use_cache=True, fill_kv_cache=True
    )

    # 2. 从噪声开始迭代去噪(10 步)
    x_t = self.sample_noise(actions_shape, device)
    dt = -1.0 / self.num_steps  # num_steps = 10

    for timestep in [1.0, 0.9, 0.8, ..., 0.1]:
        v_t = self.denoise_step(state, prefix_pad_masks,
                                past_key_values, x_t, timestep)
        x_t += dt * v_t  # Euler 步进

    return x_t  # 最终的动作预测

注意力 mask 机制

def make_att_2d_masks(pad_masks, att_masks):
    """构造 2D 注意力 mask

    图像和文本 token 之间可以互相注意(双向),
    但动作 token 使用因果注意力(只能看到之前的 token)。

    att_masks 含义:
      0 = 属于前一个因果块(可以被后面的看到)
      1 = 开始新的因果块
    """
    cumsum = torch.cumsum(att_masks, dim=1)
    att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]
    pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None]
    return att_2d_masks & pad_2d_masks

embed_suffix - 后缀编码(状态+动作+时间)

def embed_suffix(self, state, noisy_actions, timestep):
    """编码后缀"""
    # 动作嵌入
    action_emb = self.action_in_proj(noisy_actions)

    if not self.pi05_enabled:
        # PI0: 状态投影 + 时间与动作拼接
        state_emb = self.state_proj(state)
        time_emb = create_sinusoidal_pos_embedding(timestep, ...)
        action_time_emb = torch.cat([action_emb, time_emb], dim=2)
        action_time_emb = self.action_time_mlp_in(action_time_emb)
        adarms_cond = None
    else:
        # PI0.5: 时间通过 AdaRMS 条件化
        time_emb = self.time_mlp_in(time_emb)
        adarms_cond = F.silu(self.time_mlp_out(F.silu(time_emb)))

    return embs, pad_masks, att_masks, adarms_cond

核心类/函数列表

名称 类型 说明
PI0Model 类 PI0 核心模型
forward 方法 训练时单步去噪
sample_actions 方法 推理时完整采样
denoise_step 方法 带 KV cache 的单步去噪
embed_prefix 方法 编码图像+文本前缀
embed_suffix 方法 编码状态+动作+时间后缀
make_att_2d_masks 函数 构造 2D 注意力 mask
create_sinusoidal_pos_embedding 函数 正弦位置编码

与其他模块的关系

  • 使用 PaliGemmaWithExpertModel(paligemma_with_expert.py)做双流 Transformer
  • 被 PI0ForActionPrediction(modeling_pi0_torch.py)封装调用

小结

PI0Model 实现了 Flow Matching 的核心推理逻辑。关键设计是: 1. 双流架构:图像/文本走 PaliGemma 主流,动作走 Expert 支流 2. KV Cache:推理时缓存前缀(图像+文本),只重复计算后缀(动作去噪) 3. 10 步去噪:从高斯噪声迭代到最终动作,比 DDPM 的 1000 步高效很多