Loading...
文件结构总览
modeling_pi05.py 约 1305 行,分为四大块:
| 块 | 行范围 | 内容 | 核心类/函数 |
|---|---|---|---|
| A. 工具函数 | 1–280 | 兼容性、位置编码、注意力掩码、图像缩放、梯度检查点 wrapper | create_sinusoidal_pos_embedding, make_att_2d_masks, resize_with_pad_torch, compute_layer_complete |
| B. 双塔模型 | 280–460 | PaliGemma + Gemma Expert 的联合模型,处理三种前向模式 | PaliGemmaWithExpertModel |
| C. 核心 PI05 模型 | 460–920 | 训练/推理流程:加噪、去噪、嵌入、采样 | PI05Pytorch |
| D. Policy 接口 | 920–1305 | LeRobot 封装:权重加载、图像预处理、批量推理、训练 loss | PI05Policy |
推理流程总览
| 步骤 | 操作描述 | 核心函数 | 输入 | 输出 |
|---|---|---|---|---|
| 1. 图像预处理 | 多视角图像 resize+pad → 归一化到 [-1,1],缺失视角用 -1 填充 | _preprocess_images() | batch[obs_images] (B,C,H,W) | images: list[Tensor], img_masks: list[Tensor] |
| 2. Prefix 嵌入 | 图像经 SigLIP → Projector;语言 token 经 Embedding → 拼接为 prefix 序列 | embed_prefix() → embed_image() / embed_language_tokens() | images, tokens, masks | prefix_embs (B,N,D), pad_masks, att_masks |
| 3. Prefix 预填充 | PaliGemma 对 prefix 序列做双向全注意力前向,缓存所有层的 KV Cache | paligemma_with_expert.forward() (Prefix-Only) | prefix_embs, att_2d_masks_4d | past_key_values (DynamicCache) |
| 4. 噪声初始化 | 从 N(0,I) 采样初始噪声动作 | sample_noise() | (B, chunk_size, max_action_dim) | noise: Tensor (fp32) |
| 5. 去噪循环 | N 步 Euler 积分,将噪声逐步去噪为干净动作 | (循环体内) | ||
| 5a. 时间步 | t 从 1.0 线性递减到 0.0(Flow Matching 时间方向) | — | step ∈ [0,N) | time ∈ [1,0] |
| 5b. Suffix 嵌入 | 时间步经正弦编码 → 2层 MLP → AdaRMS 条件;动作经 action_in_proj 投影 | embed_suffix() | x_t, timestep | suffix_embs, adarms_cond |
| 5c. 单步去噪 | Suffix attend to Prefix KV Cache(只读),经 Gemma Expert 各层 → 预测速度场 | denoise_step() | x_t, timestep, past_key_values | v_t (B, chunk, act_dim) |
| 5d. Euler 步进 | x_{t+Δt} = x_t + Δt · v_t | — | x_t, v_t, dt=-1/N | x_t (更新后) |
| 6. 动作输出 | 截断 max_action_dim → 真实动作维度 | — | x_t (B, chunk, max_dim) | actions (B, chunk, real_dim) |
推理数据流图
点击左侧流程图节点即可跳转到对应函数详解。
flowchart TD
A["观测图像<br/><i>[list[Tensor(B,C,H,W)]]</i>"] -->|"_preprocess_images()"| B["预处理图像<br/><i>[list[Tensor(B,C,H,W)]]</i>"]
C["语言指令<br/><i>[Tensor(B,seq)]</i>"] -->|"embed_language_tokens()"| D["语言嵌入<br/><i>[Tensor(B,seq,D)]</i>"]
B -->|"embed_image()<br/>SigLIP+Projector"| E["图像嵌入<br/><i>[Tensor(B,N_img,D)]</i>"]
D --> F["Prefix 预填充<br/><i>PaliGemma (Prefix-Only)</i>"]
E --> F
F -->|"past_key_values"| G["KV Cache<br/><i>DynamicCache</i>"]
H["随机噪声<br/><i>N(0,I)</i>"] -->|"sample_noise()"| I["x_t<br/><i>[Tensor(B,chunk,max_dim)] fp32</i>"]
J["timestep<br/><i>float ∈ [1,0]</i>"] -->|"embed_suffix()<br/>sin-cos + MLP"| K["adarms_cond<br/><i>[Tensor(B,D)]</i>"]
I -->|"action_in_proj"| K
K --> L["denoise_step()<br/><i>Gemma Expert + Prefix KV Cache</i>"]
G --> L
J --> L
L -->|"v_t [B,chunk,max_dim]"| M["Euler 步进<br/><i>x_t += dt · v_t</i>"]
M -->|"× N 循环"| I
M -->|"final x_t (≈ clean action)"| N["解填充<br/><i>截断到 real_dim</i>"]
N --> O["输出动作<br/><i>[Tensor(B,chunk,real_act_dim)]</i>"]| 编号 | 节点 | 负责函数 | 点击跳转 |
|---|---|---|---|
| ①② | 观测图像 → 预处理图像 | _preprocess_images() + resize_with_pad_torch() | 点击左侧 A/B 节点 |
| ③④⑤ | 图像+语言 → Prefix 嵌入 | embed_prefix() → embed_image() / embed_language_tokens() | 点击左侧 B/C/D/E 节点 |
| ⑥ | Prefix 预填充 → KV Cache | PaliGemmaWithExpertModel.forward() (Prefix-Only) | 点击左侧 F/G 节点 |
| ⑦ | 随机噪声 → x_t | sample_noise() | 点击左侧 H/I 节点 |
| ⑧⑨ | timestep + x_t → suffix 嵌入 + AdaRMS | embed_suffix() → create_sinusoidal_pos_embedding() → time_mlp | 点击左侧 J/K 节点 |
| ⑩ | 单步去噪:Suffix attend to Prefix KV | denoise_step() → clone_past_key_values() | 点击左侧 L 节点 |
| ⑪ | Euler 步进 | 循环体内 x_t = x_t + dt * v_t | 点击左侧 M 节点 |
| ⑫⑬ | 解填充 → 最终动作 | 截断 [:, :, :real_dim] | 点击左侧 N/O 节点 |
训练流程总览
| 步骤 | 操作描述 | 核心函数 | 输入 | 输出 |
|---|---|---|---|---|
| 1. 加噪 | x_t = t·ε + (1-t)·a(直线插值) | PI05Pytorch.forward() | actions, noise, time | x_t |
| 2. 目标速度场 | u_t = ε − a(噪声减干净动作) | — | noise, actions | u_t |
| 3. Prefix 嵌入 | 图像 SigLIP + 语言 Embedding → 拼接 | embed_prefix() | images, tokens, masks | prefix_embs, pad_masks, att_masks |
| 4. Suffix 嵌入 | 时间正弦编码+2层MLP → AdaRMS 条件;动作线性投影 → suffix token | embed_suffix() | x_t, time | suffix_embs, adarms_cond |
| 5. 联合前向 | Prefix + Suffix 逐层联合 attention (compute_layer_complete × 18) | PaliGemmaWithExpertModel.forward() (联合模式) | prefix_embs, suffix_embs, adarms_cond | suffix_out |
| 6. 速度场预测 | suffix_out 截断 + action_out_proj → v_t | action_out_proj() | suffix_out[:, -chunk:] | v_t |
| 7. MSE Loss | L = ‖u_t − v_t‖² | F.mse_loss(reduction="none") | u_t, v_t | loss (B, chunk, act_dim) |
训练数据流图
flowchart TD
A["干净动作 a<br/><i>[B,chunk,act_dim]</i>"] -->|"x_t = t·ε + (1-t)·a"| B["加噪动作 x_t<br/><i>[B,chunk,act_dim]</i>"]
C["噪声 ε<br/><i>N(0,I)</i>"] --> B
D["时间 t<br/><i>Beta采样</i>"] --> B
A -->|"u_t = ε - a"| E["目标速度场 u_t<br/><i>[B,chunk,act_dim]</i>"]
C --> E
F["Prefix 嵌入<br/><i>图像+语言</i>"] --> G["联合 Attention<br/><i>compute_layer_complete × 18层</i>"]
B -->|"embed_suffix()"| H["Suffix 嵌入<br/><i>动作+时间</i>"]
H --> G
D -->|"adarms_cond"| G
G -->|"suffix_out"| I["action_out_proj"]
I -->|"v_t"| J["MSE Loss<br/><i>L = ‖u_t - v_t‖²</i>"]
E --> J| 编号 | 节点 | 负责函数 | 点击跳转 |
|---|---|---|---|
| ① | 干净动作+噪声+时间 → x_t + u_t | PI05Pytorch.forward() 前半段 | 点击左侧 A/B/C/D/E 节点 |
| ② | 图像+语言 → Prefix 嵌入 | embed_prefix() | 点击左侧 F 节点 |
| ③ | x_t+时间 → Suffix 嵌入 | embed_suffix() | 点击左侧 H 节点 |
| ④ | Prefix+Suffix 联合 attention | compute_layer_complete() + PaliGemmaWithExpertModel.forward() | 点击左侧 G 节点 |
| ⑤ | suffix_out → v_t | action_out_proj | 点击左侧 I 节点 |
| ⑥ | u_t vs v_t 比较 | F.mse_loss | 点击左侧 J 节点 |
A. 工具函数
A1. ActionSelectKwargs — 类型标注
| |
RTC(Real-Time Chunking)推理时传入的额外参数。total=False 表示所有字段可选:
inference_delay:推理延时步数,用于控制动作执行的时间窗口;prev_chunk_left_over:前一个 chunk 的剩余动作(跨 chunk 拼接);execution_horizon:执行时域长度,限制预测动作被实际执行的步数。
这些参数仅在 RTC 模式启用时被消费,普通推理会直接忽略。RCT 模式下,模型每次只预测一个 chunk 的动作序列,剩余动作会被缓存到 prev_chunk_left_over,在下一次推理时作为 prefix 继续使用。
A2. get_safe_dtype() — 跨平台 dtype 兼容
| |
确保 dtype 在特定设备上可用:
- MPS(Apple Silicon):不支持 float64(双精度),自动降级为 float32;
- CPU:不支持 bfloat16(PyTorch CPU 没有 bf16 内核),降级为 float32;float64 则保留原样(CPU 双精度性能可接受);
- 其余情况直接返回目标 dtype。
这是一个防御性适配,主要服务于 create_sinusoidal_pos_embedding 的中间计算(需要 float64 精度做 torch.linspace)。
A3. create_sinusoidal_pos_embedding() — 正弦-余弦时间编码(~L65)
| |
这是时间步到 Transformer 维度 embedding 的核心编码函数。逐行分析:
输入校验:
timeshape 必须是(B,)一维,dimension必须是偶数(sin/cos 成对);频率生成(对数均匀分布):
1 2fraction = torch.linspace(0.0, 1.0, dimension // 2) # [0, 1/(D/2-1), 2/(D/2-1), ..., 1] period = min_period * (max_period / min_period) ** fraction # 从 min_period 到 max_periodfraction在 [0,1] 之间均匀采样,period在对数空间中从min_period指数增长到max_period。例如min_period=2, max_period=10000时,period 序列为[2, 2.3, 2.7, ..., 10000]——低频到高频全覆盖。外积计算:
1 2scaling_factor = 1.0 / period * 2 * math.pi # [D/2] 频率向量 ω_j sin_input = scaling_factor[None, :] * time[:, None] # [B, 1] × [1, D/2] → [B, D/2]对 batch 中每个标量时间 $t_i$,乘以每个频率 $\omega_j$,得到相位矩阵 $\phi_{ij} = t_i \cdot \omega_j$。
拼接 sin/cos:
$$e(t) = [\sin(\omega_1 t), \cos(\omega_1 t), \sin(\omega_2 t), \cos(\omega_2 t), \dots]$$[sin(ϕ), cos(ϕ)]→ shape[B, D]。最终每个时间步 $t$ 被编码为一个 D 维向量:
为什么用对数均匀频率? 与 Transformer 的 RoPE 原理类似——低频编码长期依赖(大 period),高频编码短期精细变化,对数尺度保证频域覆盖均匀。
A4. sample_beta() — Beta 分布时间采样(~L80)
| |
标准 Beta 分布采样。时序说明:
- MPS fallback:Beta 分布的
_sample_dirichlet在 Apple Silicon 上未实现,源码注释建议 CPU 采样后搬回。此处通过.to(device)隐式处理; - 为何用 Beta 而非均匀分布? Beta(α, β) 可以控制采样密度偏向 [0,1] 的哪一端。例如 α < β 时密度偏向左侧(低时间步/低噪声阶段),让模型更多训练在"精细去噪"阶段。PI05 配置中
time_sampling_beta_alpha和time_sampling_beta_beta正是控制这一偏置。
A5. make_att_2d_masks() — 构造二维注意力掩码(~L90)
| |
这是灵活构造多种注意力模式(因果 / prefix-lm / block-causal)的核心函数。
核心算法:每个 token 有一个 mask_ar 值(来自 att_masks),规则是——token i 可以 attend to token j 当且仅当 cumsum_att_masks[j] ≤ cumsum_att_masks[i] 且两者均非 padding。
具体例子(来自 big_vision 注释):
att_masks | 含义 |
|---|---|
[1,1,1,1,1,1] | 纯因果注意力:cumsum=[1,2,3,4,5,6],对角及以下全可见 |
[0,0,0,1,1,1] | Prefix-LM:前 3 个 token 互相可见(cumsum 均为 0,0≤0 为 True);后 3 个因果 + 可见前缀 |
[1,0,1,0,1,0,0,1,0,0] | Block-Causal:每 2 个 token 组成一个 block,block 内全连接,block 间因果 |
实现细节:
| |
广播 [B, 1, N] <= [B, N, 1] 得到 [B, N, N] 的 bool 矩阵。
在 PI05 中的使用场景:
- Prefix:
att_masks全 0 →cumsum全 0 → 矩阵全 True → Prefix 内部全注意力; - Suffix(推理时):第一个 token
att_mask=1,其余att_mask=0→ 第一个 token 可见 Prefix + 自己,后续 token 因果链。
A6. clone_past_key_values() — KV Cache 深拷贝(~L115)
| |
去噪循环中每一次 denoise_step 需要独立的 KV Cache 副本。原因:
past_key_values存储 Prefix 预填充后的所有层 K/V;- 在
denoise_step中,Suffix 序列的 attention 会追加新的 K/V 到缓存(修改past_key_values内部状态); - 如果下一轮去噪复用同一个缓存对象,会将前一轮的 Suffix K/V 也一并 attend 到,造成信息泄露和序列长度递增。
因此每轮深拷贝一份干净的 Prefix KV Cache,确保 Suffix 每步都只能看到 Prefix + 自己。sliding_window 是 Gemma 2B 滑动窗口注意力的配置,直接透传。
A7. pad_vector() — 动作维度填充(~L125)
| |
简单地对最后一维右侧补零。用于将不同机器人(不同动作维度)统一 pad 到 max_action_dim。设计选择 F.pad 而非 nn.Linear 升维是因为补零不引入学习参数,且 action_out_proj 的最终线性层已经处理了维度映射。
A8. resize_with_pad_torch() — 无失真图像缩放(~L130)
| |
这个函数以保持宽高比的方式缩放图像,不足部分用 0(黑色)填充。
逐行逻辑:
通道格式检测:
1 2 3if images.shape[-1] <= 4: # 最后一维 ≤ 4 → channels-last [H,W,C] channels_last = True images = images.permute(0, 3, 1, 2) # → [B, C, H, W]启发式判据:RGB/RGBA 图像的通道数 ≤ 4,而 channels-first 格式的最后一维(Width)几乎肯定 > 4。
等比缩放:
1 2 3ratio = max(cur_width / width, cur_height / height) # 取宽高中较大的缩放比 resized_height = int(cur_height / ratio) resized_width = int(cur_width / ratio)max确保缩放后没有任何一边超过目标尺寸(短边会被 pad,长边刚好匹配)。dtype 特定裁剪:
1 2if images.dtype == torch.uint8: resized_images = torch.round(resized_images).clamp(0, 255).to(torch.uint8) elif images.dtype == torch.float32: resized_images = resized_images.clamp(0.0, 1.0)uint8 图像需要
round()取整,float32 图像直接裁剪。注意此时图像仍在 [0,1](或 [0,255])范围内,尚未归一化到 [-1,1](那个步骤在_preprocess_images中完成)。居中填充:
1 2 3pad_h0, remainder_h = divmod(height - resized_height, 2) pad_h1 = pad_h0 + remainder_h # 奇数差时右边/下边多 1px padded_images = F.pad(resized_images, (pad_w0, pad_w1, pad_h0, pad_h1), mode="constant", value=0)divmod保证居中(上左优先,多余像素给下右)。恢复通道格式:若输入是 channels-last,permute 回去。
A9. compute_layer_complete() — 联合层计算(核心)(~L195)
| |
这是 PI05 最核心的计算单元——Prefix 和 Suffix 对应的 Transformer 层在同一个 attention 矩阵中联合计算。
函数签名说明:
inputs_embeds:[prefix_embs, suffix_embs],两者的 hidden states(在不同层之间传递时会更新);layers:(paligemma_layer_i, gemma_expert_layer_i),两者的第 i 个 DecoderLayer;adarms_cond:[None, time_emb],Prefix 的 AdaRMS 条件为 None,Suffix 的 AdaRMS 条件为时间嵌入;rotary_emb:共享的 RoPE 模块(取 PaliGemma 的 rotary_emb)。
阶段 1:并行 QKV 投影 + AdaRMS(~L200-215)
| |
逐行分析:
layernorm_forward(layer.input_layernorm, hidden_states, adarms_cond[i]):- 对 Prefix 层(
adarms_cond=None):等价于标准 RMSNorm →hidden_states * weight; - 对 Suffix 层(
adarms_cond=time_emb):AdaRMS——权重和偏置由时间嵌入动态调制: $$\text{AdaRMS}(h, t) = \gamma(t) \cdot \frac{h}{\text{RMS}(h)} + \beta(t)$$ 其中 $\gamma(t) = W_\gamma \cdot \text{SiLU}(t) + b_\gamma$(从adarms_cond线性投影得到)。gate返回值用于后续 gated residual。
- 对 Prefix 层(
view(hidden_shape).transpose(1, 2):将[B, seq, num_heads * head_dim]重整为[B, num_heads, seq, head_dim]。注意这里head_dim来自各自 layer 的配置。并行处理:Prefix 和 Suffix 的 QKV 在各自的线性层执行——参数独立,仅结构对称。
阶段 2:拼接 + 联合 RoPE(~L218-232)
| |
Prefix 和 Suffix 的 QKV 在序列维度拼接,构成一个统一的 [B, H, total_len, D] 张量。
| |
dummy_tensor的 shape[B, total_seq_len, head_dim],仅用于触发rotary_emb的 shape 推断(Gemma 的 rotary_emb 不关心 content,只取 dim);position_ids是 prefix 和 suffix 拼接后的全局位置 id(后续在PI05Pytorch.forward中由torch.cumsum(pad_masks, dim=1) - 1生成);- 联合 RoPE 确保 Prefix 和 Suffix 在同一个位置编码空间中——Prefix 的图像 token 占位置 0
255,语言 token 续接位置 256511,Suffix 续接位置 512+。
阶段 3:联合 Attention(~L235-242)
| |
eager_attention_forward执行 $\text{Softmax}\left(\frac{QK^T}{\sqrt{d}} + \text{mask}\right)V$;attention_mask的 shape 是[B, 1, total_len, total_len](_prepare_attention_masks_4d已将 2D mask 升维),控制 Prefix 不可 attend to Suffix;- 注意
scaling取自 PaliGemma 层——两塔共享同一个 attention scale(实际都是 $1/\sqrt{\text{head\\_dim}}$,但明确从 PaliGemma 取是防御性写法); reshape(batch_size, -1, 1 * 8 * head_dim):1*8*head_dim是num_heads * head_dim的硬编码,此处假设 PaliGemma 和 Expert 都是 8 头。这是一个已知的局限性——如果使用非 8 头配置会出错。
阶段 4:分别 O 投影 + MLP + Gated Residual(~L244-262)
| |
Attention 输出在序列维度拆分回 Prefix 和 Suffix 部分,各自通过独立的 O 投影。注意精度的显式转换:如果 attention 输出是 fp32 但 O 投影权重是 bf16,先转换再乘法。
| |
Gated Residual(门控残差)替代了标准 x = x + sublayer(x):
其中 $g$ 由 AdaRMS 的 layernorm_forward 返回(是 sigmoid 激活后的门控值)。gate 控制残差连接的力度——$g=1$ 时完全使用 sublayer 输出,$g=0$ 时保持输入不变。这是 Pi0/PI05 对标准 Transformer 的关键改动,允许模型学习"跳过"某些层。
| |
MLP 结构:gate_proj(x) * up_proj(x) → down_proj(Gemma 的 GeGLU 变体,激活函数 gelu_pytorch_tanh)。
两个残差连接分别围绕 Attention block 和 MLP block,且各自有独立的 gate(分别来自 input_layernorm 和 post_attention_layernorm 的 AdaRMS 输出)。
返回 outputs_embeds(两塔更新后的 hidden states),作为下一层的 inputs_embeds 输入。
阶段 5:梯度检查点包装
在 PaliGemmaWithExpertModel.forward() 中,compute_layer_complete 被 torch.utils.checkpoint.checkpoint() 包裹:
| |
use_reentrant=False 使用 PyTorch 新版非重入检查点(更安全),preserve_rng_state=False 跳过 RNG 保存以减少开销(对推理过程无影响,训练中可能略微影响 dropout 但通常可接受)。这允许在 18 层 × 2 塔的配置下大幅节省显存。
B. 核心模型类详解
B1. GemmaConfig — 纯数据类(~L268)
| |
与 HuggingFace 的 GemmaConfig 不同,这是 PI05 自用的简化配置类。两种预定义变体:
| 参数 | gemma_300m | gemma_2b |
|---|---|---|
width | 1024 | 2048 |
depth | 18 | 18 |
mlp_dim | 4096 | 16384 |
num_heads | 8 | 8 |
num_kv_heads | 1 | 1 |
head_dim | 256 | 256 |
关键点:两塔都是 GQA(Grouped-Query Attention),num_kv_heads=1 意味着每个 attention 层只有 1 个 KV 头、8 个 Q 头——极大降低了 KV Cache 显存(推理时 Prefix KV Cache 约占总显存的 30-40%,GQA 将其压缩 8 倍)。
B2. PaliGemmaWithExpertModel.__init__() — 双塔初始化(~L280)
| |
参数说明:
| 参数 | 默认 | 含义 |
|---|---|---|
vlm_config | — | PaliGemma 的 GemmaConfig(Gemma 的语言部分) |
action_expert_config | — | Expert 的 GemmaConfig |
use_adarms | [False, True] | [Prefix是否用AdaRMS, Suffix是否用AdaRMS] |
precision | "bfloat16" | 整体精度,视觉路径强制 fp32 |
freeze_vision_encoder | False | 冻结 SigLIP 视觉编码器 |
train_expert_only | False | 仅训练 Expert 塔(冻结整个 PaliGemma) |
HuggingFace Config 构造:
| |
将 GemmaConfig 的字段一一映射到 HF 的 PaliGemmaConfig。adarms_cond_dim = width 表明 AdaRMS 的条件向量与 hidden_size 同维度(来自 MLP 后的时间嵌入)。
| |
Expert 塔也构造了 HF 兼容的 config,但 Suffix 塔不使用 PaliGemma 而直接使用纯 Gemma(因为不需要视觉编码器)。
| |
Expert 的 embed_tokens 置为 None——因为 Suffix 嵌入由 action_in_proj 生成(动作空间投影),不需要查字典。PiGemmaForCausalLM 的前向在 embed_tokens=None 时应直接接受 inputs_embeds。
B3. to_bfloat16_for_selected_params() — 混合精度策略
| |
设计逻辑:
- 先全局设为 bf16(节省 50% 显存/带宽);
- 再将视觉路径和所有归一化层恢复为 fp32。
为什么视觉路径必须 fp32?
- SigLIP 的 patch embedding 和 projector 涉及
trunc_normal初始化的小数值,bf16 的动态范围不足(最小正数 ~9.2e-41 vs fp32 的 ~1.4e-45)会导致梯度消失; - 注释中明确写道 “never toggle”——训练时如果视觉路径在 fp32 和 bf16 之间切换,PyTorch 优化器会报 “same dtype” 错误。
为什么归一化层(RMSNorm)也保持 fp32?
- RMSNorm 计算 $\frac{x}{\sqrt{\frac{1}{d}\sum x_i^2}}$,其中分母 RMS 值可能极小(尤其在深层的残差流中),bf16 精度不足会导致 NaN。
B4. _set_requires_grad() 和 train() — 冻结策略
| |
train(mode) 的覆盖很关键:super().train(mode) 会将所有子模块设回 training 模式,但冻结部分需要额外强制 eval()——确保 BN/Dropout 在冻结模块中不激活。
B5. embed_image() 和 embed_language_tokens()
| |
- 输入 image 强制转为 fp32 → SigLIP 前向 →
pooler_output(即经过 multi-modal projector 后的[B, N_patches, width]特征); - 输出转回输入 dtype(如果外部是 bf16,将 fp32 特征转回 bf16),节省后续计算带宽。
| |
标准 token embedding 查表,返回 [B, seq_len, width]。
B6. PaliGemmaWithExpertModel.forward() — 三种前向模式
这是双塔模型的调度中心。根据 inputs_embeds 的内容分为三种模式:
模式 1:Prefix-Only(推理时预填充)
| |
- 触发条件:
inputs_embeds = [prefix_embs, None]; - 仅走 PaliGemma 的语言模型,不使用 Expert;
use_cache=True返回 DynamicCache,后续推理复用;adarms_cond[0] = None(PI05 中 Prefix 塔use_adarms=False)。
模式 2:Suffix-Only(推理时去噪步骤)
| |
- 触发条件:
inputs_embeds = [None, suffix_embs]; - 仅走 Gemma Expert,通过
past_key_values携带 Prefix 的 KV Cache → Suffix 在 attention 中可以看到 Prefix 的所有 token; use_cache=False:每轮去噪后不需要缓存 Suffix 的 KV(被clone_past_key_values刷新);adarms_cond[1] = time_emb:时间嵌入注入 Expert 的 AdaRMS。
模式 3:联合模式(训练)
| |
- 触发条件:
inputs_embeds = [prefix_embs, suffix_embs]两者都非 None; - 18 层(gemma_300m)或 18 层(gemma_2b)逐层交替计算;
compute_layer_complete在上面已详细分析——每层先拼接两塔的 hidden states 做联合 attention,再分别 MLP;- 最终 Norm 也使用 AdaRMS(Prefix 塔的
adarms_cond=None退化为普通 RMSNorm,Suffix 塔的adarms_cond=time_emb正常注入); prefix_past_key_values = None:训练时不缓存 KV Cache。
返回值统一格式:([prefix_output, suffix_output], prefix_past_key_values)
B7. PI05Pytorch.__init__() — 核心模型初始化
| |
use_adarms=[False, True] 是 PI05 区别于 PI0 的核心标志。PI0 使用 FiLM(use_adarms=[True, True]),PI05 将 Prefix 塔的 FiLM 替换为普通 RMSNorm,仅 Suffix 塔保留 AdaRMS。
| |
四组可学习投影:
| 投影 | 维度变换 | 作用 |
|---|---|---|
action_in_proj | max_action_dim → D | 将动作向量嵌入到 Expert 的 hidden 空间 |
action_out_proj | D → max_action_dim | 将 Expert 输出映射回动作空间(预测速度场) |
time_mlp_in | D → D | 时间嵌入 MLP 第一层 |
time_mlp_out | D → D | 时间嵌入 MLP 第二层 |
注意:PI05 没有 state_proj——这与 PI0 不同。PI0 将机器人状态(关节角度等)通过单独的投影并入,PI05 将状态信息视为动作维度的一部分统一处理。
| |
torch.set_float32_matmul_precision("high") 允许 PyTorch 使用 TF32 tensor cores(A100/H100 上),加速 fp32 矩阵乘法约 2-3×。torch.compile 对训练前向和推理函数做 JIT 编译。
B8. gradient_checkpointing_enable/disable / _apply_checkpoint
| |
三层梯度检查点:Prefix 语言模型、视觉塔、Expert 模型。配合 compute_layer_complete 中的 HF checkpoint wrapper,训练时每层只保存输入,反向传播时重计算。
| |
训练时启用检查点的统一接口。非训练或未启用时直接调用。
B9. _prepare_attention_masks_4d() — 2D→4D 掩码转换
| |
att_2d_masks是 bool 矩阵[B, N, N](True=可见,False=不可见);- 升维到
[B, 1, N, N]适配多头 attention; OPENPI_ATTENTION_MASK_VALUE是一个极小值(通常 ~-2.381e38 或torch.finfo(torch.float32).min),加到 attention logits 上使 Softmax 后不可见位置的权重变为 0。
B10. sample_noise() / sample_time()
| |
标准高斯噪声。shape 是 (B, chunk_size, max_action_dim),强制 fp32 保证去噪数值精度。
| |
时间从 Beta 分布采样后线性变换:
$$t = t_{\text{beta}} \cdot \text{scale} + \text{offset}$$典型配置(来自 openpi):alpha=1.5, beta=1.0, scale=1.0, offset=0.0 → 采样偏向较小时刻(更多训练在低噪声精细去噪阶段)。scale < 1 时限制 t ∈ [offset, offset+scale]。
B11. embed_prefix() — Prefix 嵌入拼接
| |
img_mask 是 [B] 的 bool(0=缺失视角,1=有效视角),扩展为 [B, N_img]。att_masks 全部填 0 → 所有图像 token 互相可见。
| |
语言 token 的 att_masks 同样为 0 → Prefix 内部所有 token(图像 + 语言)共享同一个 attention group,实现双向全注意力。
| |
B12. embed_suffix() — Suffix 嵌入(动作+时间)
| |
时间编码的维度 = Expert 的 hidden size(不是 action dim),因为时间嵌入需要和 hidden states 交互(作为 AdaRMS 条件和残差流相加)。
| |
时间嵌入经过 time_mlp_in → SiLU → time_mlp_out → SiLU 的双层 MLP。第二个 SiLU 是关键的——它使 adarms_cond 始终非负(SiLU 最小值 ≈ -0.278),限制 AdaRMS 门控的动态范围。
| |
关键设计: Suffix 的 embedding 仅包含 action_emb(不拼接 time_emb)。时间信息完全通过 AdaRMS 注入每一层——这比直接拼接更优雅,因为时间信息以乘法形式调制归一化参数,而非作为 token 参与 attention。
att_masks = [1, 0, 0, ..., 0](chunk_size 个元素):
- 第一个 token 的
att_mask=1→cumsum从 1 开始 → 可以 attend 到cumsum ≤ 1的所有前缀 token; - 后续 token
att_mask=0→cumsum保持 1 → 因果注意力(每个 token 可 attend 到前缀 + 当前 chunk 内之前的所有 token)。
B13. PI05Pytorch.forward() — 训练前向
| |
Flow Matching 定义:从干净动作 $a$ 到噪声 $\epsilon$ 的直线路径:
$$x_t = t \cdot \epsilon + (1-t) \cdot a, \quad u_t = \epsilon - a$$- $t=0$:$x_0 = a$(干净动作),$u_0 = \epsilon - a$(从干净指向噪声);
- $t=1$:$x_1 = \epsilon$(纯噪声),$u_1 = \epsilon - a$(不变)。
速度场 $u_t$ 在整个路径上是常数——这正是 Flow Matching 的 Straight Line Flow 性质。
| |
检查第一层 Q 投影权重的 dtype:如果是 bf16,将 embeddings 也转为 bf16(减少计算开销)。如果视觉路径保持 fp32(to_bfloat16_for_selected_params 的副作用),此处不转换。
| |
position_ids 通过累积有效 token 数生成:padding 位置的 cumsum 不变(mask=0 → 不加),保证位置 id 连续。
| |
past_key_values=None+use_cache=False:训练时不缓存 KV;inputs_embeds=[prefix_embs, suffix_embs]:触发联合模式;[:, -chunk_size:]:从输出中截取 Suffix 部分——模型输出包含 prefix 和 suffix token,只取最后chunk_size个;- 预测前转回 fp32(loss 计算需要高精度)。
| |
返回未归约的 MSE——由上层 PI05Policy.forward() 根据 reduction 参数决定是否取均值。
B14. PI05Pytorch.sample_actions() — 推理采样(核心)
| |
阶段 1:Prefix 预填充
| |
关键细节:
_attn_implementation = "eager":torch.compile 默认使用sdpa/flash_attention_2,但这些实现与DynamicCache的某些路径不兼容(特别是use_cache=True时的 KV 追加),强制回退到 eager;inputs_embeds=[prefix_embs, None]:触发 Prefix-Only 模式,仅计算 Prefix KV Cache;use_cache=True:返回past_key_values(DynamicCache)。
阶段 2:去噪循环
| |
时间线性递减:从 t=1(纯噪声)到 t=0(干净动作),步长 $\Delta t = -1/N$。
| |
闭包捕获 past_key_values 和 prefix_pad_masks,暴露简洁的 (x_t) → v_t 接口。这个包装是为了 RTC 兼容——RTC 需要调用 original_denoise_step_partial 作为 fallback。
| |
RTC 模式下,rtc_processor.denoise_step 可能在 chunk 边界上拼接前一个 chunk 的剩余动作,或在执行时域限制下截断预测——这是一个实时控制的扩展功能,普通推理直接走 else 分支。
| |
调试模式记录每一步的中间态,用于可视化去噪过程。
B15. PI05Pytorch.denoise_step() — 单步去噪
| |
full_att_2d_masks 的 shape 是 [B, suffix_len, prefix_len + suffix_len]:
- 左半部分
[B, suffix_len, prefix_len]:来自prefix_pad_masks的广播——Suffix 的每个 token 可以看到 Prefix 的所有有效(padding=1)token; - 右半部分
[B, suffix_len, suffix_len]:来自make_att_2d_masks的 Suffix 因果注意力。
| |
Suffix 的位置 id 从 Prefix 的总 token 数开始递增,保证全局位置编码连续。
| |
每次 denoise_step 都深拷贝 Prefix KV Cache——防止 Suffix 的新 KV 污染 Prefix 缓存(分析见 A6)。
| |
C. Policy 接口类详解
C1. PI05Policy.__init__() — Policy 初始化
| |
标准 LeRobot Policy 初始化协议。reset() 初始化动作队列(用于 n_action_steps > 1 时的 action chunk 调度)。
C2. PI05Policy.from_pretrained() — 权重加载
| |
步骤 A:Key Fix(_fix_pytorch_state_dict_keys)
| |
处理 checkpoint 与当前代码键名不一致的多种情况(详见下文 C3)。
步骤 B:Remap Prefix
| |
所有不含 model. 前缀的 key 都加上——因为 PI05Policy 把 PI05Pytorch 存为 self.model,HF 的 save_pretrained 不会自动加前缀,但 LeRobot 的加载需要。
| |
C3. _fix_pytorch_state_dict_keys() — 键名兼容性修复
处理 OpenPI checkpoint 到 PI05 的多种键名差异:
Case 1: AdaRMS 不兼容 → 跳过
| |
如果 checkpoint 中的 Expert 层是 .weight 格式(普通 RMSNorm),但当前模型使用 AdaRMS(use_adarms=True),则跳过——因为 AdaRMS 的权重结构不同(W_gamma + W_beta vs 单个 weight)。同样处理 final_norm。
Case 2: MLP 命名差异
| |
PI05 将 PI0 中的 action_time_mlp_* 重命名为 time_mlp_*。
Case 3: state_proj 不存在
| |
PI05 没有 state_proj——如果 checkpoint 包含(来自 PI0),直接丢弃。
Case 4: lm_head.weight → embed_tokens.weight
| |
PaliGemma 的 lm_head 权重与 embed_tokens 共享(tied weights),checkpoint 可能只存了一份。这里将其复制到 embedding 位置(推理不需要 lm_head)。
C4. _preprocess_images() — 图像预处理
| |
有效视角处理
| |
img.shape[1] == 3 是 heuristic 判据——假设只有 channels-first 格式的第二维是 3(C=3),但这也可能是 H=3 的极小图像。LeRobot 数据集几乎都是 [B,C,H,W],风险可忽略。
缺失视角处理
| |
全 -1 并非非法输入:SigLIP 归一化后 [-1, 1] 范围中,-1 对应原始 [0,1] 范围中的 0(纯黑)。模型训练时应包含缺失视角的数据增强(随机 dropout 相机),因此推理时遇到缺失视角可以合理处理。
C5. prepare_action() / select_action() / predict_action_chunk()
| |
训练时动作维度 pad 到 max_action_dim。
| |
Action Chunk 调度策略:
- 模型预测 50 步动作,取前 10 步(
n_action_steps)入队; - 每次 select_action 弹出 1 步;
- 队列空后重预测——这意味着每 10 步触发一次模型推理,大幅降低推理频率(50 步才需要 5 次推理)。
| |
C6. PI05Policy.forward() — 训练 Loss
| |
reduction="none" 返回 per-sample loss 用于 RA-BC(Reward-Augmented Behavioral Cloning)等需要逐样本权重的训练方法。
C7. _get_default_peft_targets() — PEFT/LoRA 默认目标
| |
LoRA 默认只微调 Expert 塔的 Q/V 投影 + 输入输出投影,Prefix 塔保持冻结。
总结:关键数据流
训练:
图像 + 语言 + 动作 → embed_prefix + embed_suffix
→ 联合 attention (compute_layer_complete × 18层)
→ suffix_out → action_out_proj → v_t
→ MSE(u_t, v_t)
推理 (sample_actions):
图像 + 语言 → embed_prefix → Prefix KV Cache (预填充)
→ noise = N(0,I); x_t = noise
→ 循环 N 步:
x_t, timestep → embed_suffix → denoise_step (Suffix attend to Prefix KV)
→ action_out_proj → v_t
→ x_t = x_t + dt * v_t
→ 返回 x_t (≈ 干净动作)