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

修改强化学习代码后张量规格一致但训练停滞问题求助

DQN改CNN后训练停滞的修复方案

核心问题分析

你修改的代码存在两个关键错误,导致训练无法推进:

  1. 手动循环破坏计算图连续性:逐个处理batch样本并手动构建tensor,不仅效率低下,还切断了模型前向传播的计算图链路,梯度无法正确回传。
  2. 目标网络误用:计算next_state_values时错误使用了policy_net而非target_net,违背了DQN双网络分离的核心设计,导致训练信号混乱。

修复步骤

1. 修复state_action_values计算

删除手动循环的代码,恢复批量处理逻辑(确保state_batch的形状符合CNN输入要求,比如(batch_size, channels, height, width)):

# 直接用批量前向传播+gather,和原逻辑一致,适配CNN输出
state_action_values = self.policy_net(state_batch).gather(1, action_batch)
print(state_action_values.size())  # torch.Size([128, 1])

2. 修复next_state_values计算

将误用的policy_net改回target_net,同时保持批量处理逻辑:

next_state_values = torch.zeros(self.batch_size, device=self.device)
with torch.no_grad():
    # 用target_net计算,批量处理而非循环
    next_state_values[non_final_mask] = self.target_net(non_final_next_states).max(1)[0]
print(next_state_values.size())  # torch.Size([128])

额外注意事项

  • 确保CNN网络的输出维度和原全连接网络一致:比如原全连接输出(batch_size, num_actions),CNN最后一层也要输出相同形状的张量。
  • 不要手动设置requires_grad=True:模型前向输出的张量会自动追踪梯度,手动设置会干扰PyTorch的自动微分机制。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 17:39:59