1. 从 infer.py 到 sample_actions:pi05 策略推理全链路解析
在强化学习和机器人控制领域,策略推理的完整实现链路往往隐藏在层层封装之下。本文将以 pi05 策略模型为例,拆解从入口脚本到神经网络前向传播的全过程。这个案例特别适合想要深入理解现代机器人策略部署细节的开发者——我们将看到 JAX 与 PyTorch 混合编程的实践,以及如何高效处理多模态输入。
提示:本文假设读者已具备基础的强化学习知识,熟悉策略网络、观测空间等概念。若遇到陌生术语,建议先查阅相关基础资料。
1.1 入口脚本:infer.py 的核心使命
infer.py 作为整个推理流程的入口点,承担着三个关键职责:
-
配置加载与策略初始化
通过命令行参数接收配置名和检查点路径,这是工业级代码的典型做法——允许在不修改代码的情况下切换不同实验配置。例如:bash复制
python infer.py libero_pi05 /path/to/checkpoints对应的配置解析逻辑会加载
libero_pi05对应的 YAML 或 JSON 文件,其中包含:- 观测空间定义(图像分辨率、关节状态维度等)
- 策略网络超参数(Transformer 层数、注意力头数等)
- 预处理/后处理参数(归一化范围、动作缩放系数等)
-
策略对象的动态构建
create_trained_policy是一个工厂方法,其内部完成:- 从检查点恢复模型参数(可能涉及 JAX 的
flax.serialization或 PyTorch 的torch.load) - 实例化完整的策略流水线(包含预处理模块、神经网络、后处理模块)
- 将模型设置为评估模式(关闭 dropout 等随机性操作)
- 从检查点恢复模型参数(可能涉及 JAX 的
-
观测封装与推理触发
原始观测需要被封装成策略期望的格式。对于 LIBERO 基准任务,典型的观测字典可能包含:python复制obs = { 'rgb': np.ndarray # (H, W, 3) 的摄像头图像 'depth': np.ndarray # (H, W) 的深度图 'proprio': np.ndarray # (7,) 的机械臂关节状态 }
1.2 策略推理的核心链路
当调用 policy.infer(obs) 时,实际触发的处理流程可分为五个阶段:
阶段一:观测预处理
python复制def preprocess(obs):
# 图像标准化 (假设原始像素值范围[0,255])
obs['rgb'] = obs['rgb'].astype(np.float32) / 127.5 - 1.0
# 深度图归一化 (根据传感器量程调整)
obs['depth'] = (obs['depth'] - DEPTH_MEAN) / DEPTH_STD
# 关节状态缩放
obs['proprio'] = scale_proprio(obs['proprio'])
# 增加批次维度
return {k: v[None] for k, v in obs.items()}
阶段二:神经网络前向传播
这是最复杂的部分,涉及:
- 多模态特征的编码(CNN 处理图像,MLP 处理关节状态)
- 跨模态特征的融合(通过 Transformer 的交叉注意力)
- 动作序列的生成(自回归采样或一步预测)
阶段三:动作后处理
包括:
- 去除批次维度
- 反归一化(将网络输出的 [-1,1] 映射到实际动作范围)
- 加入安全限制(速度限幅、关节角度限制等)
阶段四:缓存更新(适用于自回归策略)
如果策略使用 Tra
