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

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)

额外优化建议

  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)
  1. 清理冗余配置:代码中注册了自定义环境learning_rate_env,但最终使用的是CartPole-v1,建议删除未用到的环境注册代码,避免混淆。

  2. 简化GPU设置:num_gpus_per_worker=1/1可直接写为num_gpus_per_worker=1,效果一致。

后续LSTM策略迁移提示

切换到LSTM模型时需注意:

  • 在forward方法中正确处理seq_lens参数,维护序列维度
  • 确保每次调用forward时传递正确的模型隐状态(state)
  • 若使用自定义LSTM层,需显式处理时间维度与批量维度的顺序

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 11:52:02