paligemma_with_expert.py — 实现了 PI0 的双流 Transformer 架构¶
文件路径:
verl/experimental/vla/models/pi0_torch/model/paligemma_with_expert.py模块路径:verl.experimental.vla.models.pi0_torch.model.paligemma_with_expert
文件概述¶
实现了 PI0 的双流 Transformer 架构。核心思想是在标准 PaliGemma 模型旁边添加一个 Expert 流,两者共享注意力计算但有独立的 MLP 和层归一化。这是 PI0 模型的"大脑"。
架构概览¶
PaliGemma 流(主流) Expert 流(动作流)
图像 + 文本嵌入 状态 + 动作嵌入
↓ ↓
[LayerNorm] [AdaRMS Norm]
↓ ↓
┌────────── 共享注意力 ──────────┐
│ Q_pg, K_pg, V_pg Q_ex, K_ex, V_ex │
│ ↓ ↓ │
│ ┌─── concat K, V ───┐ │
│ │ Scaled Dot-Product Attention │ │
│ └────────────────────────────────┘ │
└──── O_pg ──────── O_ex ──────────────┘
↓ ↓
[MLP_pg] [MLP_ex]
↓ ↓
PG 输出 Expert 输出
核心类¶
GemmaRMSNorm - 支持 AdaRMS 的归一化¶
class GemmaRMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6, use_ada_rms_norm=False):
if use_ada_rms_norm:
# PI0.5: 使用时间步条件化的 AdaRMS
self.dense = nn.Linear(dim, dim * 3) # 输出 scale, shift, gate
else:
self.weight = nn.Parameter(torch.zeros(dim))
def forward(self, x, cond=None):
normed = self._norm(x.float())
if self.use_ada_rms_norm:
scale, shift, gate = torch.chunk(self.dense(cond), 3, dim=-1)
normed = normed * (1 + scale) + shift
return normed, gate # gate 用于残差连接
return normed * (1 + self.weight)
GemmaAttentionWithExpert - 双流共享注意力¶
class GemmaAttentionWithExpert(nn.Module):
"""两个流共享注意力计算,但有独立的 QKV 投影"""
def __init__(self, layer_idx, ...):
# PaliGemma 流的投影
self.q_proj = nn.ModuleList([
nn.Linear(paligemma_hidden_size, ...), # PG
nn.Linear(expert_hidden_size, ...), # Expert
])
# K, V, O 投影类似
# 共享 RoPE
self.rope_embedding = RoPEEmbedding(dim=head_dim)
def forward(self, inputs_embeds, position_ids, attention_mask, ...):
# 分别投影两个流的 Q, K, V
q_pg = self.q_proj[0](inputs_embeds[0])
q_ex = self.q_proj[1](inputs_embeds[1])
# 拼接进行联合注意力
q = torch.cat([q_pg, q_ex], dim=1)
k = torch.cat([k_pg, k_ex], dim=1)
v = torch.cat([v_pg, v_ex], dim=1)
# 计算注意力
att_output = F.scaled_dot_product_attention(q, k, v, attn_mask=...)
# 分割输出
out_pg = self.o_proj[0](att_output[:, :pg_len])
out_ex = self.o_proj[1](att_output[:, pg_len:])
return [out_pg, out_ex]
GemmaDecoderLayerWithExpert - 双流解码层¶
class GemmaDecoderLayerWithExpert(nn.Module):
"""一个完整的双流解码层"""
def forward(self, inputs_embeds, adarms_cond, ...):
# 1. 输入层归一化(Expert 流可用 AdaRMS)
for i, hidden_states in enumerate(inputs_embeds):
if self.pi05_enabled and adarms_cond[i] is not None:
normed, gate = self.input_layernorms[i](hidden_states, adarms_cond[i])
else:
normed = self.input_layernorms[i](hidden_states)
# 2. 共享注意力
attn_outputs = self.self_attn(normed_embeds, ...)
# 3. 门控残差连接
after_attn = self.gated_residual(residual, attn_output, gate)
# 4. 独立 MLP
mlp_out = self.mlps[i](post_norm)
# 5. 再次门控残差
output = self.gated_residual(residual, mlp_out, mlp_gate)
PaliGemmaWithExpertModel - 完整的双流模型¶
class PaliGemmaWithExpertModel(nn.Module):
def __init__(self, pi05_enabled=False):
# SigLIP 视觉编码器
self.vision_tower = SiglipVisionTransformer(siglip_config)
# 视觉到语言投影
self.multi_modal_projector = PaliGemmaMultiModalProjector(...)
# 语言嵌入
self.embed_tokens = nn.Embedding(vocab_size, hidden_size)
# 18 层双流解码器
self.layers = nn.ModuleList([
GemmaDecoderLayerWithExpert(i, pi05_enabled, ...)
for i in range(18)
])
RoPEEmbedding - 旋转位置编码¶
class RoPEEmbedding(nn.Module):
"""预计算的 RoPE 嵌入,提高效率"""
def __init__(self, dim, max_wavelength=10000, max_seq_len=8192):
# 预计算 sin/cos 值
inv_freq = 1.0 / (max_wavelength ** (2.0/dim * torch.arange(dim//2)))
positions = torch.arange(max_seq_len)
freqs = torch.outer(positions, inv_freq)
self.register_buffer("cos_cached", torch.cos(freqs))
self.register_buffer("sin_cached", torch.sin(freqs))
核心类列表¶
| 名称 | 说明 |
|---|---|
GemmaRMSNorm |
支持 AdaRMS 的 RMS 归一化 |
SiglipVisionTransformer |
SigLIP 视觉编码器 |
PaliGemmaMultiModalProjector |
视觉投影器 |
RoPEEmbedding |
旋转位置编码 |
GemmaAttentionWithExpert |
双流共享注意力 |
GemmaMLP |
门控 MLP |
GemmaDecoderLayerWithExpert |
双流解码层 |
PaliGemmaWithExpertModel |
完整双流模型 |
与其他模块的关系¶
- 被
PI0Model(modeling_pi0.py)使用 - 使用 HuggingFace transformers 的 SigLIP 组件
小结¶
这个文件实现了 PI0 的核心创新:双流 Transformer。关键设计是让视觉/语言流和动作流共享注意力计算(使得动作可以"看到"视觉和语言信息),同时保持独立的 MLP 和归一化(使得两个流可以有不同的表示空间)。PI0.5 通过 AdaRMS 归一化引入时间步条件化,是对 PI0 的改进。