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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 15:54:53