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

PyTorch实现DQN时MSELoss报target与input维度不匹配如何解决

DQN训练时MSE损失维度不匹配报错解决

问题根源

报错核心是q_target计算过程中张量维度不匹配:

  • q_eval是经过gather提取的对应动作Q值,形状为(batch_size, 1)即(32,1),这部分逻辑是正确的
  • 计算q_next.max(1)[0]时,PyTorch默认会压缩被取最大值的维度,返回的批量最大Q值形状为(batch_size,)即(32,);但你的奖励张量b_r形状是(32,1),两个张量相加时触发广播机制,最终得到形状为(32,32)的错误q_target,和(32,1)的q_eval输入MSE损失时就会触发维度不匹配警告。

修改方法

只需要在计算q_next行最大值时,添加keepdim=True参数,保持输出维度为(32,1),和b_r维度对齐即可。

需修改的代码段

原learn()函数中q_target计算代码:

q_next = self.target_net(b_s_).detach()  
q_target = b_r + gamma * q_next.max(1)[0]

修改为:

q_next = self.target_net(b_s_).detach()  
q_target = b_r + gamma * q_next.max(1, keepdim=True)[0]

修改后q_target形状为(32,1),和q_eval维度完全匹配,警告会直接消失。

额外优化建议

  • 你的经验回放池容量为100、batch_size为32,建议增加判断:当self.memory_cntr >= batch_size时再启动学习,否则会采样到大量初始全0的无意义数据,影响训练收敛速度
  • choose_action函数中用到的epsilon和stringlist变量目前未在类内定义,需要补充对应声明,否则运行时会触发变量未定义错误
  • 网络输出维度设置为n_action=32,请确认和实际任务的离散动作空间大小一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.02 08:48:25