DQN 2015 Nature 论文复现:Atari Pong 游戏 84x84 像素输入实战(附 PyTorch 代码)

发布时间:2026/8/30 6:59:03

DQN 2015 Nature 论文复现:Atari Pong 游戏 84x84 像素输入实战(附 PyTorch 代码)
DQN 2015 Nature 论文复现Atari Pong 游戏 84x84 像素输入实战附 PyTorch 代码当DeepMind在2015年首次提出DQN算法并在Nature上发表时整个强化学习领域为之震动。这项研究首次证明一个单一的深度强化学习智能体能够在数十款Atari 2600游戏中达到人类水平的表现。本文将带您从零开始使用现代PyTorch框架完整复现这一里程碑式的工作特别聚焦于Atari Pong游戏的实现细节。1. 环境配置与预处理在开始构建DQN之前我们需要先搭建适合的训练环境。Atari游戏的原始输入为210×160像素的RGB图像这对计算资源提出了较高要求。遵循原论文的方法我们将进行以下预处理import gym import numpy as np from collections import deque import torch import torch.nn as nn import torch.optim as optim class AtariPreprocessor: def __init__(self, env_name, frame_skip4, history_length4): self.env gym.make(env_name) self.frame_skip frame_skip self.history deque(maxlenhistory_length) def reset(self): frame self.env.reset() processed self._process_frame(frame) for _ in range(self.history.maxlen): self.history.append(processed) return np.stack(self.history) def step(self, action): total_reward 0.0 for _ in range(self.frame_skip): frame, reward, done, info self.env.step(action) total_reward reward if done: break processed self._process_frame(frame) self.history.append(processed) return np.stack(self.history), total_reward, done, info def _process_frame(self, frame): # 转换为灰度图并调整大小 frame frame.mean(axis2) # RGB转灰度 frame frame[34:34160, :160] # 裁剪得分区域 frame frame[::2, ::2] # 下采样到80x80 return frame.astype(np.float32) / 255.0关键预处理步骤包括帧堆叠将连续的4帧堆叠作为网络输入提供时序信息帧跳过每4帧执行一次动作提高训练效率图像裁剪移除不相关的屏幕区域如得分显示灰度转换将RGB三通道简化为单通道归一化将像素值缩放到[0,1]范围注意原论文使用84×84分辨率但实际实现中80×80也是常见选择。确保测试时与训练分辨率一致。2. DQN网络架构设计DQN的核心是一个深度卷积神经网络其架构设计直接影响了特征提取能力。以下是PyTorch实现class DQN(nn.Module): def __init__(self, action_dim): super(DQN, self).__init__() self.conv1 nn.Conv2d(4, 32, kernel_size8, stride4) self.conv2 nn.Conv2d(32, 64, kernel_size4, stride2) self.conv3 nn.Conv2d(64, 64, kernel_size3, stride1) self.fc1 nn.Linear(7*7*64, 512) self.fc2 nn.Linear(512, action_dim) def forward(self, x): x x.float() / 255.0 # 确保输入归一化 x torch.relu(self.conv1(x)) x torch.relu(self.conv2(x)) x torch.relu(self.conv3(x)) x x.view(x.size(0), -1) # 展平 x torch.relu(self.fc1(x)) return self.fc2(x)网络结构参数对比如下层类型参数输出尺寸激活函数卷积层132个8×8滤波器步长420×20×32ReLU卷积层264个4×4滤波器步长29×9×64ReLU卷积层364个3×3滤波器步长17×7×64ReLU全连接层1512单元512ReLU输出层动作空间维度action_dim线性3. 经验回放与目标网络DQN的两个关键创新点需要特别实现class ReplayBuffer: def __init__(self, capacity): self.buffer deque(maxlencapacity) def push(self, state, action, reward, next_state, done): self.buffer.append((state, action, reward, next_state, done)) def sample(self, batch_size): indices np.random.choice(len(self.buffer), batch_size, replaceFalse) states, actions, rewards, next_states, dones zip(*[self.buffer[idx] for idx in indices]) return ( torch.FloatTensor(np.array(states)), torch.LongTensor(np.array(actions)), torch.FloatTensor(np.array(rewards)), torch.FloatTensor(np.array(next_states)), torch.FloatTensor(np.array(dones)) ) def __len__(self): return len(self.buffer) class DQNAgent: def __init__(self, action_dim, lr1e-4, gamma0.99, tau1e-3): self.policy_net DQN(action_dim) self.target_net DQN(action_dim) self.target_net.load_state_dict(self.policy_net.state_dict()) self.optimizer optim.Adam(self.policy_net.parameters(), lrlr) self.gamma gamma self.tau tau def update_target(self): # 软更新目标网络 for target_param, policy_param in zip(self.target_net.parameters(), self.policy_net.parameters()): target_param.data.copy_( self.tau * policy_param.data (1.0 - self.tau) * target_param.data ) def get_action(self, state, epsilon): if np.random.random() epsilon: return np.random.randint(self.policy_net.fc2.out_features) with torch.no_grad(): q_values self.policy_net(state.unsqueeze(0)) return q_values.argmax().item()经验回放和目标网络的作用经验回放打破数据相关性提高样本效率目标网络稳定训练过程防止Q值过高估计软更新缓慢更新目标网络参数τ通常取0.0014. 完整训练流程将上述组件整合为完整的训练系统def train_dqn(env_namePongNoFrameskip-v4, batch_size32, buffer_size100000, total_steps1000000, learning_starts10000, target_update1000, gamma0.99, epsilon_start1.0, epsilon_end0.1, epsilon_decay100000): env AtariPreprocessor(env_name) agent DQNAgent(env.env.action_space.n) buffer ReplayBuffer(buffer_size) state env.reset() episode_reward 0 total_rewards [] epsilon epsilon_start for step in range(1, total_steps 1): # ε-贪心策略 epsilon epsilon_end (epsilon_start - epsilon_end) * \ np.exp(-1. * step / epsilon_decay) # 选择并执行动作 action agent.get_action(torch.FloatTensor(state), epsilon) next_state, reward, done, _ env.step(action) episode_reward reward # 存储转移样本 buffer.push(state, action, reward, next_state, done) state next_state # 训练阶段 if len(buffer) learning_starts and step % 4 0: batch buffer.sample(batch_size) states, actions, rewards, next_states, dones batch # 计算当前Q值 current_q agent.policy_net(states).gather(1, actions.unsqueeze(1)) # 计算目标Q值 with torch.no_grad(): next_q agent.target_net(next_states).max(1)[0] target_q rewards (1 - dones) * gamma * next_q # 计算损失并更新 loss nn.MSELoss()(current_q.squeeze(), target_q) agent.optimizer.zero_grad() loss.backward() agent.optimizer.step() # 更新目标网络 if step % target_update 0: agent.update_target() # 回合结束处理 if done: total_rewards.append(episode_reward) print(fStep: {step}, Reward: {episode_reward}, Epsilon: {epsilon:.2f}) state env.reset() episode_reward 0 return total_rewards训练过程中的关键参数设置参数推荐值作用batch_size32每次更新的样本数量buffer_size100,000经验回放缓存大小learning_starts10,000开始学习前的随机探索步数target_update1,000目标网络更新频率gamma0.99未来奖励折扣因子epsilon_start1.0初始探索率epsilon_end0.1最终探索率epsilon_decay100,000探索率衰减步数5. 训练技巧与性能优化在实际训练中以下几个技巧可以显著提升性能奖励裁剪将正奖励设为1负奖励设为-1有助于不同游戏间的泛化reward np.clip(reward, -1, 1)帧差分处理取连续帧的最大值消除Atari游戏的闪烁效果frame np.maximum(frame, last_frame)梯度裁剪防止梯度爆炸稳定训练过程for param in agent.policy_net.parameters(): param.grad.data.clamp_(-1, 1)学习率调度随着训练进展降低学习率scheduler optim.lr_scheduler.StepLR(agent.optimizer, step_size250000, gamma0.1)在Pong游戏中典型的训练曲线会经历以下阶段随机探索期0-10k步智能体随机移动胜率约50%初步学习期10k-100k步开始学习基本击球策略策略优化期100k-500k步发展出位置控制和反击策略稳定表现期500k步达到人类水平胜率超过90%6. 结果评估与可视化训练完成后我们需要评估智能体的实际表现def evaluate(agent, env, episodes10): total_rewards [] for _ in range(episodes): state env.reset() episode_reward 0 done False while not done: action agent.get_action(torch.FloatTensor(state), epsilon0.05) state, reward, done, _ env.step(action) episode_reward reward env.render() # 可视化游戏过程 total_rewards.append(episode_reward) return np.mean(total_rewards)对于Pong游戏成功的训练应能达到以下指标指标预期值说明平均奖励18每局21分制达到人类水平胜率90%对阵内置AI的获胜概率训练时间8-12小时使用现代GPU如RTX 30807. 进阶改进方向原始DQN虽然强大但仍有改进空间。以下是几个值得尝试的扩展Double DQN减少Q值高估问题next_actions agent.policy_net(next_states).max(1)[1] next_q agent.target_net(next_states).gather(1, next_actions.unsqueeze(1))优先经验回放更高效地利用重要样本td_error (current_q - target_q).abs() priority (td_error 1e-5).pow(alpha)Dueling架构分离状态价值和优势函数class DuelingDQN(nn.Module): def __init__(self, action_dim): super().__init__() # 共享特征提取层 self.feature nn.Sequential(...) # 价值流 self.value nn.Linear(512, 1) # 优势流 self.advantage nn.Linear(512, action_dim) def forward(self, x): features self.feature(x) value self.value(features) advantage self.advantage(features) return value advantage - advantage.mean()在实际项目中我发现使用Dueling架构能显著提升Pong游戏的训练速度通常在200k步左右就能达到不错的表现。而优先回放则在更复杂的游戏中效果更为明显。

相关新闻

机器学习数据集划分实战:6:2:2 黄金比例与 10 折交叉验证的 5 个关键抉择

机器学习数据集划分实战:6:2:2 黄金比例与 10 折交叉验证的 5 个关键抉择

2026/8/29 21:13:21

机器学习数据集划分实战:6:2:2黄金比例与10折交叉验证的5个关键抉择 当你在深夜调试一个图像识别模型时,验证集上的准确率突然从92%暴跌到65%,而训练集指标却依然稳步上升——这不是恐怖故事的开头,而是每个机器学习工程师都可能遇…

朴素贝叶斯分类器 Python 实现:从零手写 2 个核心函数与拉普拉斯平滑

朴素贝叶斯分类器 Python 实现:从零手写 2 个核心函数与拉普拉斯平滑

2026/8/28 20:58:19

从零实现朴素贝叶斯分类器:核心函数与平滑技术实战1. 朴素贝叶斯算法原理精要朴素贝叶斯分类器是基于贝叶斯定理与特征条件独立假设的分类方法。其核心思想是通过先验概率和条件概率来计算后验概率,从而实现对样本的分类决策。让我们先看一个简单的例子&…

动态规划算法 Python 实现:从 4 阶段图例到 100x100 栅格地图路径规划

动态规划算法 Python 实现:从 4 阶段图例到 100x100 栅格地图路径规划

2026/8/29 6:46:44

动态规划算法 Python 实现:从 4 阶段图例到 100x100 栅格地图路径规划在机器人导航和游戏开发中,路径规划是一个核心问题。想象一下,你正在开发一个仓库物流机器人,它需要在复杂的货架迷宫中找到最优路径搬运货物。传统的暴力搜索…

字节跳动前端面试全流程复盘:从简历到Offer的实战经验

字节跳动前端面试全流程复盘:从简历到Offer的实战经验

2026/8/30 6:51:38

记一次字节跳动前端面试经历:从简历投递到Offer的全过程复盘今年年中我经历了一次完整的字节跳动前端岗位面试流程,从最初的简历筛选到最终的技术定级,前后持续了将近三周。整体感受是:字节的面试节奏快、考察维度全面、非常关注候…

【关注可白嫖源码】--课程设计--毕业设计--智慧养老服务系统[编号:project64990](案件分析)

【关注可白嫖源码】--课程设计--毕业设计--智慧养老服务系统[编号:project64990](案件分析)

2026/8/30 6:51:38

摘 要目前社会老龄化进程不断加快,传统养老机构人员调配、业务处理等大多依靠人工操作,存在信息传递慢、护理记录容易出错、费用核算繁杂等问题,不能适应规模化、精细化的管理要求。因此设计并实现一个整合资源、优化流程的养老服务系统有重…

Replit云开发实战:从零搭建并部署Web应用

Replit云开发实战:从零搭建并部署Web应用

2026/8/30 6:51:38

现在,越来越多的开发者开始把“写代码”这件事从本地搬回浏览器。不是因为本地环境不好,而是因为现代云开发平台确实解决了很多真实工程痛点——尤其是当你同时拥有多台电脑、需要快速验证想法、或者想和小伙伴一起协作时,这种“打开浏览器就…

R语言进阶——众数回归模型(modalreg)

R语言进阶——众数回归模型(modalreg)

2026/8/30 6:51:38

目录0、引言1、核心思想1.1、目标函数的构成1.2、与均值和中位数的关系2、R语言包——modalreg2.1、包的简介2.2、快速上手安装安装与加载2.3、使用案例生成模拟数据(内置工具)训练模型查看结果与预测2.4、关键参数说明💡 使用建议⚠️ 注意事…

语音算法工程师笔试全解析:从信号处理到深度学习考点梳理

语音算法工程师笔试全解析:从信号处理到深度学习考点梳理

2026/8/30 6:51:38

2018年秋招,语音算法岗还没有现在这么内卷,但竞争已经相当激烈。欢聚时代那会儿手里有YY和虎牙这两个大流量产品,语音算法工程师要从音频采集处理一直管到内容理解,笔试出得相当硬核。我当年参加过B卷的考试,后来又在音…

欢聚时代校招笔试真题解析:产品、数据、运营、市场四岗位答题方法论

欢聚时代校招笔试真题解析:产品、数据、运营、市场四岗位答题方法论

2026/8/30 6:41:38

2018年那一场在成都的欢聚时代校招笔试,我到现在还有印象。当时帮团队筛选了不少产品岗和运营岗的试卷,卷子上那些题目看似杂,其实背后全是一个套路:大厂想在校招阶段就看清楚你的思维方式,而不仅仅是知识储备。这份A卷…

备战数据库管理工程师校招:索引、事务、备份恢复核心考点解析

备战数据库管理工程师校招:索引、事务、备份恢复核心考点解析

2026/8/30 0:01:07

每年校招季我都会接触不少准备数据库方向笔试的同学,看到最多的状态就是:简历上写着“熟悉 MySQL”“了解索引优化”,一碰到数据库管理工程师的笔试卷,却在索引、事务、锁、备份恢复这些题目上翻车。网易这套 2018 校园招聘数据库…

数字电路时序基石:深入理解建立时间与保持时间

数字电路时序基石:深入理解建立时间与保持时间

2026/8/30 0:01:07

1. 这不是“背公式”的事:时间参数到底在约束什么你翻过数字电路教材,一定见过这两个词:建立时间(Setup Time)和保持时间(Hold Time)。它们常被并列写在触发器(Flip-Flop&#xff09…

蓝桥杯国赛超声波测距机:从单片机原理到嵌入式系统实战

蓝桥杯国赛超声波测距机:从单片机原理到嵌入式系统实战

2026/8/30 0:01:07

1. 项目缘起:从赛题到超声波测距机的诞生第八届蓝桥杯单片机设计与开发国赛的题目,我至今记忆犹新。它没有直接给出一个花哨的名字,而是用“超声波测距机”这个朴实无华的功能描述,精准地勾勒出了考核的核心。对于当时备赛的我而言…

备战数据库管理工程师校招:索引、事务、备份恢复核心考点解析

备战数据库管理工程师校招:索引、事务、备份恢复核心考点解析

2026/8/30 0:01:07

每年校招季我都会接触不少准备数据库方向笔试的同学,看到最多的状态就是:简历上写着“熟悉 MySQL”“了解索引优化”,一碰到数据库管理工程师的笔试卷,却在索引、事务、锁、备份恢复这些题目上翻车。网易这套 2018 校园招聘数据库…

数字电路时序基石:深入理解建立时间与保持时间

数字电路时序基石:深入理解建立时间与保持时间

2026/8/30 0:01:07

1. 这不是“背公式”的事:时间参数到底在约束什么你翻过数字电路教材,一定见过这两个词:建立时间(Setup Time)和保持时间(Hold Time)。它们常被并列写在触发器(Flip-Flop&#xff09…

蓝桥杯国赛超声波测距机:从单片机原理到嵌入式系统实战

蓝桥杯国赛超声波测距机:从单片机原理到嵌入式系统实战

2026/8/30 0:01:07

1. 项目缘起:从赛题到超声波测距机的诞生第八届蓝桥杯单片机设计与开发国赛的题目,我至今记忆犹新。它没有直接给出一个花哨的名字,而是用“超声波测距机”这个朴实无华的功能描述,精准地勾勒出了考核的核心。对于当时备赛的我而言…

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

2026/8/28 7:35:26

一天写完毕业论文在2026年已不再是天方夜谭。2026年最炸裂、实测能大幅提速的AI论文写作工具,覆盖选题构思、文献整理、内容生成、格式排版等核心场景,真正帮你高效搞定论文难题。 一、全流程王者:一站式搞定论文全链路(一天定稿首…

导师推荐!2026最新AI论文工具测评与实用推荐

导师推荐!2026最新AI论文工具测评与实用推荐

2026/8/28 7:34:51

2026年真正好用的AI论文工具,核心看生成的论文质量、低AI味、格式正确、学术适配四大指标。综合实测,千笔AI、ThouPen、豆包、DeepSeek、Grammarly 是当前最值得推荐的梯队,覆盖从免费到付费、从中文到英文、从文科到理工的全场景需求。 一、…

告别游戏崩溃:XCOM 2模组管理器的智能革命

告别游戏崩溃:XCOM 2模组管理器的智能革命

2026/8/28 7:34:35

告别游戏崩溃:XCOM 2模组管理器的智能革命 【免费下载链接】xcom2-launcher The Alternative Mod Launcher (AML) is a replacement for the default game launchers from XCOM 2 and XCOM Chimera Squad. 项目地址: https://gitcode.com/gh_mirrors/xc/xcom2-lau…