跳转至

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 的改进。