Ray RLLIB自定义MLP策略触发AssertionError的排查求助
问题排查与修复
错误根源
核心问题是值函数输出形状与奖励形状不匹配:
- 你的自定义模型中,critic层输出的张量形状为
(batch_size, 1) - GAE计算环节期望的奖励张量形状是
(batch_size,) - 两者维度不一致触发了
assert vf_preds.shape == rewards.shape断言失败
代码修复
修改CustomMLPModel中的value_function方法,挤压掉输出张量的最后一维:
@override(ModelV2) def value_function(self): return self._value_out.squeeze(dim=-1)
额外优化建议
- 简化初始化逻辑:用
super().__init__()替代分开调用父类构造函数,代码更简洁:
def __init__(self, obs_space, action_space, num_outputs, model_config, name): super().__init__(obs_space, action_space, num_outputs, model_config, name) self.fc1 = nn.Linear(obs_space.shape[0], 128) self.fc2 = nn.Linear(128, 128) self.actor = nn.Linear(128, num_outputs) self.critic = nn.Linear(128, 1)
清理冗余配置:代码中注册了自定义环境
learning_rate_env,但最终使用的是CartPole-v1,建议删除未用到的环境注册代码,避免混淆。简化GPU设置:
num_gpus_per_worker=1/1可直接写为num_gpus_per_worker=1,效果一致。
后续LSTM策略迁移提示
切换到LSTM模型时需注意:
- 在
forward方法中正确处理seq_lens参数,维护序列维度 - 确保每次调用forward时传递正确的模型隐状态(
state) - 若使用自定义LSTM层,需显式处理时间维度与批量维度的顺序
内容的提问来源于stack exchange,提问作者David
相关产品推荐
相关产品推荐

