You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

RLlib中PPO+LSTM模型LSTM层初始化维度不匹配问题

解决RLlib中PPO+LSTM维度不匹配问题

问题根源分析

从报错和调试输出来看,核心问题集中在三个点:

  • RLlib传入的LSTM状态维度顺序与PyTorch LSTM的期望顺序不匹配
  • LSTM输入张量的形状不符合batch_first=True的设置要求
  • value_function依赖的张量未正确赋值

具体修正步骤

1. 修正LSTM状态维度顺序

RLlib传入的state中,每个状态张量(h0/c0)的形状是[batch_size, num_layers, hidden_size](即调试中的[32,1,15]),但PyTorch LSTM要求的隐藏状态形状是[num_layers, batch_size, hidden_size](即[1,32,15]),需要转置前两个维度:

h0, c0 = state
# 转置前两个维度,适配PyTorch LSTM要求
h0 = h0.transpose(0, 1)
c0 = c0.transpose(0, 1)

2. 修正LSTM输入张量形状

你的LSTM设置了batch_first=True,意味着输入形状应为[batch_size, seq_len, input_size]。当前x的形状是[32,128],需要在seq_len维度(第1维)添加维度,而非第0维:

# 替换x.unsqueeze(0)为x.unsqueeze(1)
x = x.unsqueeze(1)  # 形状变为[32,1,128]
x, new_state = self.lstm(x, (h0, c0))
# 后续只需要去掉seq_len维度
x = x.squeeze(1)

3. 修正新状态的返回格式

从PyTorch LSTM返回的new_state是[num_layers, batch_size, hidden_size]格式,需要转换回RLlib期望的[batch_size, num_layers, hidden_size]:

new_h, new_c = new_state
new_state = [new_h.transpose(0, 1), new_c.transpose(0, 1)]

4. 修复value_function的未赋值问题

value_function引用的self._last_layer_out需要在forward中提前赋值:

# 在LSTM输出后添加
self._last_layer_out = x

完整修正后的代码

class CustomTorchModel(TorchModelV2, nn.Module):
    def __init__(self, obs_space, action_space, num_outputs, model_config, name):
        TorchModelV2.__init__(self, obs_space, action_space, num_outputs, model_config, name)
        nn.Module.__init__(self)

        self.obs_size = obs_space.shape[0]  # Assuming obs_space is already shaped (12,)
        self.hidden_dim = 128 # Hidden dimension for LSTM and Dense layers
        self.lstm_hidden_state_size = 15  # Size of LSTM hidden state
        
        self.input_layer = nn.Linear(self.obs_size + 1 + action_space.shape[0], self.hidden_dim)
        self.lstm = nn.LSTM(self.hidden_dim, self.lstm_hidden_state_size, batch_first=True)
        self.output_layer = nn.Linear(self.lstm_hidden_state_size, num_outputs)
        self.logits_layer = nn.Linear(self.lstm_hidden_state_size, action_space.shape[0])
        self.log_std = nn.Parameter(torch.zeros(action_space.shape[0]))
    
    @override(TorchModelV2)
    def forward(self, input_dict, state, seq_lens):
        obs = input_dict["obs"]
        prev_reward = input_dict["prev_rewards"].unsqueeze(-1)
        last_actions = input_dict["prev_actions"]

        x = torch.cat([obs, prev_reward, last_actions], dim=-1)
        x = torch.relu(self.input_layer(x))
        
        # 转换状态维度适配PyTorch LSTM
        h0, c0 = state
        h0 = h0.transpose(0, 1)
        c0 = c0.transpose(0, 1)

        # 调整输入形状适配batch_first=True
        x = x.unsqueeze(1)
        x, new_state = self.lstm(x, (h0, c0))
        x = x.squeeze(1)
        
        # 保存最后一层输出用于value计算
        self._last_layer_out = x

        # 转换新状态维度适配RLlib要求
        new_h, new_c = new_state
        new_state = [new_h.transpose(0, 1), new_c.transpose(0, 1)]

        logits = self.logits_layer(x)
        return logits, new_state

    def value_function(self):
        return self.output_layer(self._last_layer_out)
    
    @override(TorchModelV2)
    def get_initial_state(self):
        # 返回[num_layers, hidden_size],RLlib会自动扩展batch_size维度
        return [torch.zeros(1, self.lstm_hidden_state_size),
                torch.zeros(1, self.lstm_hidden_state_size)]

关键说明

  • RLlib与PyTorch对LSTM状态的维度顺序定义不同,必须在两者之间做转换
  • batch_first=True时,输入张量的第一维必须是batch_size,第二维是序列长度
  • 确保value_function依赖的张量在forward中正确赋值

内容的提问来源于stack exchange,提问作者Federica Tonti

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 03:52:10