跳转至

mlp.py — 通用的多层感知机(MLP)模块

文件路径: verl/experimental/vla/models/modules/mlp.py 模块路径: verl.experimental.vla.models.modules.mlp

文件概述

通用的多层感知机(MLP)模块,支持可配置的层数、激活函数和初始化方法。被 SAC 的 Critic 网络和其他需要 MLP 的地方使用。

关键代码

class MLP(nn.Module):
    """可配置的多层感知机

    Args:
        input_dim: 输入维度
        hidden_dims: 隐藏层维度列表,如 [1024, 512, 256]
        output_dim: 输出维度
        activation: 激活函数类型 ("relu", "gelu", "tanh" 等)
        init_method: 权重初始化方法 ("normal", "xavier", "kaiming" 等)
    """

    def __init__(self, input_dim, hidden_dims, output_dim,
                 activation="relu", init_method="normal"):
        super().__init__()
        layers = []
        prev_dim = input_dim

        # 构建隐藏层
        for hidden_dim in hidden_dims:
            layers.append(nn.Linear(prev_dim, hidden_dim))
            layers.append(self._get_activation(activation))
            prev_dim = hidden_dim

        # 输出层(不加激活函数)
        layers.append(nn.Linear(prev_dim, output_dim))

        self.network = nn.Sequential(*layers)

        # 初始化权重
        self._init_weights(init_method)

    def forward(self, x):
        return self.network(x)

使用示例

在 SAC Critic 中的使用:

critic_head = MLP(
    input_dim=2150,           # 2048(视觉特征) + 32(状态) + 70(动作)
    hidden_dims=[1024, 512, 256],
    output_dim=1,             # 输出 Q 值
    activation="relu",
    init_method="normal",
)

与其他模块的关系

  • 被 PI0ForActionPrediction(modeling_pi0_torch.py)的 SAC Critic 使用
  • 被 RobDataParallelSACActor(sac/sac_actor.py)使用

小结

通用 MLP 构建块,通过参数化配置避免了重复编写类似的网络结构。