DQN在Atari Pong环境中无法学习有效策略的问题排查
DQN在Atari Pong训练失效的问题分析与修复方案
核心问题定位
你的实现与DeepMind 2015年Nature DQN论文的关键差异是额外添加了批归一化(Batch Normalization),这是导致Atari环境下训练失效的最主要原因:
- 原论文的卷积层未使用BN,Atari游戏的输入是0-255的灰度像素,分布相对稳定;而BN会强制对每个batch的特征做归一化,经验回放的随机采样会导致batch间分布波动,破坏模型对输入分布的学习,最终引发Q值坍缩、奖励持续下降。
- 简单测试环境状态空间小,BN的负面影响不明显,但Atari环境高维度、分布复杂,BN的统计偏移会直接导致模型无法收敛。
具体修复思路与代码调整
1. 移除所有批归一化层
删除模型中所有BN层及相关调用,还原原论文的网络结构:
class NatureQN(nn.Module): def __init__(self, env, config, device): super(NatureQN, self).__init__() self.env = env self.config = config self.device = device # 移除所有BN层 self.conv1 = nn.Conv2d(in_channels=env.observation_space.shape[0], out_channels=32, kernel_size=8, stride=4, dtype=torch.float32) self.conv2 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=4, stride=2, dtype=torch.float32) self.conv3 = nn.Conv2d(in_channels=64, out_channels=64, kernel_size=3, stride=1, dtype=torch.float32) self.fc1 = nn.Linear(in_features=3136, out_features=512, dtype=torch.float32) self.fc2 = nn.Linear(in_features=512, out_features=env.action_space.n, dtype=torch.float32) self.ReLU = nn.ReLU() self.optimizer = optim.RMSprop(self.parameters(), lr=self.config.lr_begin, alpha=self.config.squared_gradient_momentum, eps=self.config.rms_eps) self.criterion = nn.SmoothL1Loss() def forward(self, x): x = self.ReLU(self.conv1(x)) x = self.ReLU(self.conv2(x)) x = self.ReLU(self.conv3(x)) if len(x.shape) == 3: x = torch.flatten(x) elif len(x.shape) == 4: x = torch.flatten(x, start_dim=1) x = self.ReLU(self.fc1(x)) x = self.fc2(x) return x
2. 还原固定学习率设置
原论文使用固定学习率0.00025,无学习率衰减。你的线性衰减会导致训练后期学习率过低,无法有效更新模型:
- 删除学习率调度器
LambdaLR的初始化代码 - 在
train_on_minibatch中移除self.approx_network.scheduler.step()调用
3. 修正RMSprop参数
原论文中RMSprop的alpha(平方梯度动量系数)为0.95,你的配置是0.9,会降低梯度更新的稳定性:
self.optimizer = optim.RMSprop(self.parameters(), lr=0.00025, alpha=0.95, eps=1e-2)
4. 确认学习频率
原论文设定每4步执行一次梯度更新,检查AtariDQNConfig中learning_freq是否设置为4。若设为1会导致梯度噪声过大,模型难以收敛。
5. 清理冗余代码
在train_on_minibatch中删除self.target_network.optimizer.zero_grad(),目标网络不参与训练,无需清零梯度。
6. 验证输入归一化逻辑
确保process_state函数将回放缓冲区存储的uint8状态转换为0-1之间的float32张量(对应原论文将灰度像素除以255的操作):
def process_state(self, state): return state.to(torch.float32) / 255.0
训练建议
- 严格遵循原论文超参数:epsilon从1.0线性衰减到0.1(500万步)、目标网络每10000步更新、batch_size=32、回放池大小100万,这些参数经过大量验证,不要随意修改。
- 训练时长:原论文中Pong需要约400-500万步才能达到人类水平,修复上述问题后建议至少训练到500万步再评估效果。
内容的提问来源于stack exchange,提问作者Rohan Patel
相关产品推荐
相关产品推荐

