修改强化学习代码后张量规格一致但训练停滞问题求助
DQN改CNN后训练停滞的修复方案
核心问题分析
你修改的代码存在两个关键错误,导致训练无法推进:
- 手动循环破坏计算图连续性:逐个处理batch样本并手动构建tensor,不仅效率低下,还切断了模型前向传播的计算图链路,梯度无法正确回传。
- 目标网络误用:计算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
相关产品推荐
相关产品推荐

