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 构建块,通过参数化配置避免了重复编写类似的网络结构。