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
相关产品推荐
相关产品推荐

