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 的优势是原生支持连续动作空间,无需离散化带来的精度损失。