gym Atari环境下DQN算法实现与调参指南
简介基于 Python 在 gym Atari 环境中实现 DQN 算法及其变体DDQN 等的课程设计资源包面向正在学习深度强化学习、需要完成类似实验或作业的高年级本科生与研究生。资源包含完整 Python 源码、训练日志、结果演示与说明文档共255个文件以 py 脚本、txt 配置与日志、png/gif 演示图、markdown 笔记及 PDF 说明为主压缩包约 5.05MB。作者记录了真实训练中遇到的设备性能限制、内存不足导致中断以及调参迭代的完整过程这些排错思路对初学者尤为实用目录中还有 Dockerfile 与依赖配置文件便于复现实验环境。目前已有204人学习下载适合作为从理论到代码实现的参考帮助读者少走弯路、快速上手 DQN 实验。1. 用 Python 在 gym Atari 里跑 DQN为什么我还要劝你先想清楚这一点很多人第一次接触 DQN 算法都是从 gym 的 Atari 环境入的门一个 210×160 的彩色画面一根操纵杆一个得分数字跑起来像在看老式街机。但等你真的把 DQN 代码跑起来会发现一个比想象中扎心的事实在 Atari 上复现 DQN 并不难难的是复现出和论文接近的效果。同样的代码换个随机种子可能就从“学会打砖块”变成“原地发呆”这个现象在强化学习圈子里被戏称为“强化学习玄学”。我今天讲的这套方案就是围绕标题里「gym Atari 环境实现 DQN 算法及其变体」展开的完整落地路径面向那些已经学过一点 Python 和深度学习、想亲手把 DQN 从理论变成能跑通、能调参、能对比的实验代码的从业者。你会看到怎么搭环境、怎么设计智能体、怎么处理 Atari 的帧栈和奖励缩放以及最重要的——当训练曲线不涨的时候该查哪些地方而不是盲目改学习率。2. 先把地基打牢gym Atari 环境的选择与安装坑2.1 用 gym 加载 Atari 环境版本不同命就不同要跑 DQN第一步是把环境装好。这里有个容易翻车的细节gym 库在 0.26 版本前后 API 变化很大而 Atari 环境的加载方式也经历了好几次调整。我常用的组合是gymnasiumgym 的继任者配合ale-py如果你还在用老的gym0.21会遇到env.unwrapped和render模式不兼容的问题。安装时建议直接指定版本pip install gymnasium[atari] pip install gymnasium[accept-rom-license]第一条命令会把 gymnasium 和 Atari 依赖装上第二条命令会同时安装 ROM 许可相关配置。注意accept-rom-license这个 extras 安装的是允许自动下载 ROM 的许可声明如果没有它某些环境会直接报No ROM is found。然后是环境创建方式。以经典的Pong为例正确写法是import gymnasium as gym env gym.make( ALE/Pong-v5, frameskip4, repeat_action_probability0.0, full_action_spaceFalse, render_modergb_array, )这里几个参数值得说明frameskip4表示每 4 帧执行一次动作这能减少计算量并且让智能体有时间观察运动趋势repeat_action_probability0.0是模拟 Atari 主机摇杆随机漂移的概率论文里一般设为 0但要复现更真实的环境也可以设为 0.25full_action_spaceFalse表示只用 6 个离散动作基本操作而不是全部 18 个动作render_modergb_array是为了在训练时拿像素作为观测而不是开窗口渲染。如果你用的是老版本 gymgym.make(PongNoFrameskip-v4)这种环境名也还能用但不建议再折腾了新项目直接上 gymnasium 更省心。2.2 Atari 观测空间为什么原始画面不能直接喂给网络Atari 的原始观测是 (210, 160, 3) 的 RGB 图像单帧画面包含大量无关信息比如背景颜色、记分牌数字。直接用原始图像做输入网络容量需求大训练也慢。常见做法是做一组预处理把画面缩成 (84, 84) 的灰度图然后堆叠最近 4 帧作为状态。这里我用一个包装器把整个预处理流程串起来import cv2 import numpy as np import gymnasium as gym from gymnasium import spaces class AtariPreprocessor(gym.Wrapper): def __init__(self, env, height84, width84, frame_stack4): super().__init__(env) self.height height self.width width self.frame_stack frame_stack self.frames np.zeros((frame_stack, height, width), dtypenp.uint8) self.observation_space spaces.Box( low0, high255, shape(frame_stack, height, width), dtypenp.uint8 ) def _preprocess(self, frame): frame cv2.cvtColor(frame, cv2.COLOR_RGB2GRAY) frame cv2.resize(frame, (self.width, self.height), interpolationcv2.INTER_AREA) return frame def reset(self, **kwargs): obs, info self.env.reset(**kwargs) frame self._preprocess(obs) self.frames np.stack([frame] * self.frame_stack, axis0) return self.frames, info def step(self, action): obs, reward, terminated, truncated, info self.env.step(action) frame self._preprocess(obs) self.frames np.concatenate([self.frames[1:], frame[None, ...]], axis0) return self.frames, reward, terminated, truncated, info这段代码的逻辑是reset时用同一帧填充 4 个堆叠槽位保证初始状态不是全黑step时每来一帧新画面就丢掉最旧的一帧把新帧拼到最后。这样网络看到的状态就包含了最近 4 帧的运动信息比如球速、球的方向、挡板移动趋势。这里有个参数细节灰度化之后不做归一化到 [0,1] 的操作而是保持 uint8 的 0-255 范围。原因是 DQN 原文采用了这种输入而且 uint8 省内存训练时候可以在网络内部再归一化。如果你要跑浮点版本记得把dtype改成float32并在送入网络前手动除以 255。2.3 奖励缩放Clipping 和 Episode Life 是两个必须理解的开关Atari 环境自带一个特性很多游戏单帧奖励是 1、0、-1但也有一些游戏奖励数值很大比如BankHeist单步奖励可能到几十。DQN 原文对所有游戏统一做了 reward clipping把奖励限制到 [-1, 1]作用是防止不同游戏的奖励量纲差异导致训练不稳定。但要注意奖励 clip 会丢失信息——比如一个游戏里 5 和 1 代表的进度完全不同clip 之后网络只能区分“正/负/零”学习效率反而可能下降。我的建议是先做最基础的 reward clip等训练稳定后可以尝试去掉 clip观察曲线变化。在代码里clip 只需要一行reward np.clip(reward, -1.0, 1.0)还有一个 Atari 特有的开关叫lose life即“丢命”。很多游戏里智能体有三条命丢一条命不代表游戏结束但丢命往往意味着“做错了大事”。常见做法是在丢命时把 episode 截断terminate让智能体更快感知到错误。gymnasium 的 Atari 环境默认不会自动这样做需要自己包装class EpisodicLifeEnv(gym.Wrapper): def __init__(self, env): super().__init__(env) self.lives 0 def reset(self, **kwargs): obs, info self.env.reset(**kwargs) self.lives self.env.unwrapped.ale.lives() return obs, info def step(self, action): obs, reward, terminated, truncated, info self.env.step(action) lives self.env.unwrapped.ale.lives() if lives self.lives: terminated True self.lives lives return obs, reward, terminated, truncated, info这个包装器会在丢命时把terminated变为 True从而结束当前 episode。注意这只影响训练时的终止逻辑不影响环境内部真实状态。如果你要对比论文效果这个开关必须打开因为 DQN 原文的评估方式就是“每个生命作为一个 episode”来统计得分的。3. DQN 及其变体的核心实现网络、经验池与更新逻辑3.1 用 PyTorch 实现 DQN 网络从经典结构到双网络DQN 的网络结构并不复杂一个输入 4 帧堆叠 (4, 84, 84) 的卷积网络。我用的结构是论文里最常见的版本两层卷积 两层全连接输出为动作数量维度。PyTorch 实现如下import torch import torch.nn as nn import torch.nn.functional as F class DQN(nn.Module): def __init__(self, input_channels4, num_actions6): super().__init__() self.conv1 nn.Conv2d(input_channels, 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(64 * 7 * 7, 512) self.fc2 nn.Linear(512, num_actions) def forward(self, x): x x.float() / 255.0 x F.relu(self.conv1(x)) x F.relu(self.conv2(x)) x F.relu(self.conv3(x)) x x.reshape(x.size(0), -1) x F.relu(self.fc1(x)) return self.fc2(x)为什么卷积核是 8、4、3步长是 4、2、1这个参数组合是从 Atari 的原始 DQN 论文里继承下来的经过大量实验验证对不同游戏基本都适用。最后一层全连接输出的是每个动作的 Q 值不做 softmax因为 Q 值是回归目标不是概率分布。计算时先把输入从 uint8 转 float 并除以 255这一步在 forward 里做能避免在预处理时多存一份浮点数据。训练时使用双网络结构一个在线网络policy_net用于选择动作和计算当前 Q 值一个目标网络target_net用于计算目标 Q 值。目标网络不参与梯度更新而是每隔一定步数从在线网络复制参数。这个机制是 DQN 稳定训练的关键没有它用同一个网络同时评估和更新会导致目标不断变化训练过程像在追自己的尾巴。3.2 经验回放池为什么必须用队列而不是列表经验回放是把每一步的 (state, action, reward, next_state, done) 存进一个缓冲池训练时随机采样一小批。这样做有两个好处一是打破样本的时间相关性二是提高样本利用率。一个容易忽略的细节是经验池的容量直接决定训练稳定性和内存占用。我用一个简单的collections.deque实现from collections import deque import random class ReplayBuffer: def __init__(self, capacity100000): 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): batch random.sample(self.buffer, batch_size) states, actions, rewards, next_states, dones zip(*batch) return ( np.array(states), np.array(actions), np.array(rewards, dtypenp.float32), np.array(next_states), np.array(dones, dtypenp.float32), ) def __len__(self): return len(self.buffer)deque(maxlencapacity)会在容量满时自动丢弃最老的样本这比手动维护列表要可靠得多。容量设置多少一般 Atari 上用 100 万但个人电脑内存可能不够一个样本是 4×84×84 的 uint8100 万样本就是 4×84×84×1e6 ≈ 28GB。所以在本地实验时建议先用 10 万容量或者把状态压缩成灰度 uint8 存储而不是 float 数组。我用 10 万容量时Pong 训练 100 万步也够用因为 Pong 的 episode 很短样本多样性充足。采样时的random.sample是无放回采样这是标准做法。如果要用优先经验回放PER那就要换成带权重采样的实现这个稍后讲到变体时会提。3.3 训练主循环epsilon 衰减、目标网络更新与时序差计算训练主循环是 DQN 的骨架我把它拆成几个模块动作选择、环境交互、经验存储、网络更新。整个流程跑起来长这样def train_dqn(env, policy_net, target_net, buffer, optimizer, args): total_steps 0 update_steps 0 epsilon args.epsilon_start for episode in range(args.max_episodes): state, _ env.reset() episode_reward 0 done False while not done: # 选择动作epsilon-greedy if random.random() epsilon: action env.action_space.sample() else: with torch.no_grad(): q_values policy_net(torch.tensor(state[None]).float()) action q_values.argmax(dim1).item() next_state, reward, terminated, truncated, _ env.step(action) done terminated or truncated reward np.clip(reward, -1.0, 1.0) buffer.push(state, action, reward, next_state, done) state next_state episode_reward reward total_steps 1 # 更新网络 if len(buffer) args.batch_size: update_steps 1 update_policy(policy_net, target_net, optimizer, buffer, args) # 周期性同步目标网络 if total_steps % args.target_update 0: target_net.load_state_dict(policy_net.state_dict()) # 衰减 epsilon epsilon max(args.epsilon_min, args.epsilon_start - total_steps * args.epsilon_decay) if episode % 10 0: print(fEpisode {episode}, reward {episode_reward:.1f}, epsilon {epsilon:.3f}) return policy_net这里epsilon从 1.0 线性衰减到 0.1前几百万步里智能体几乎完全随机探索之后逐渐转为利用已学到的策略。这个衰减速度很重要衰减太快会导致智能体过早陷入局部最优衰减太慢会浪费训练时间。update_policy函数是核心def update_policy(policy_net, target_net, optimizer, buffer, args): states, actions, rewards, next_states, dones buffer.sample(args.batch_size) states_t torch.tensor(states).float() actions_t torch.tensor(actions).long() rewards_t torch.tensor(rewards).float() next_states_t torch.tensor(next_states).float() dones_t torch.tensor(dones).float() # 当前 Q 值 q_values policy_net(states_t).gather(1, actions_t.unsqueeze(1)).squeeze(1) # 目标 Q 值用目标网络计算取 max Q with torch.no_grad(): next_q_values target_net(next_states_t).max(1)[0] targets rewards_t args.gamma * next_q_values * (1 - dones_t) loss F.mse_loss(q_values, targets) optimizer.zero_grad() loss.backward() optimizer.step()这里有几个值得展开的细节。actions_t.unsqueeze(1)是为了配合gather操作从 Q 值矩阵中取出实际采取动作对应的 Q 值。next_q_values.max(1)[0]是经典 DQN 的目标计算方式它总是取最大值这会导致 Q 值过高估计overestimation是后续 Double DQN 要解决的问题。(1 - dones_t)是因为终止状态的下一状态没有未来奖励目标 Q 值应该只等于当前奖励。如果不乘这个掩码智能体会在 episode 结束时错误地继续学习一个不存在的下一步造成价值高估。args.gamma通常设为 0.99。这个值越大智能体越在乎长期收益越小越在乎短期收益。Atari 游戏通常用 0.99因为很多得分需要连续操作才能获得。3.4 DQN 变体逐个实现Double DQN、Dueling DQN、优先经验回放标题里说“及其变体”最常见的三个变体是 Double DQN、Dueling DQN 和 Prioritized Replay。这三个改动都很小但各自解决一个具体的痛点。Double DQN解决的是 Q 值过估计问题。经典 DQN 用max_1取目标价值时如果网络存在噪声最大值总是偏向高估。Double DQN 的做法是用在线网络选择最优动作再用目标网络计算该动作的价值。这样即使在线网络高估了某个动作目标网络给出的价值也相对中立从而缓解过估计。# Double DQN 目标计算 with torch.no_grad(): next_actions policy_net(next_states_t).argmax(dim1, keepdimTrue) next_q_values target_net(next_states_t).gather(1, next_actions).squeeze(1) targets rewards_t args.gamma * next_q_values * (1 - dones_t)改动只有两行先找在线网络认为最好的动作再用目标网络取那个动作的 Q 值。经验上 Double DQN 在多数 Atari 游戏上都能让训练更稳定尤其是 Pong、Breakout 这种奖励稀疏的游戏。Dueling DQN改变的是网络结构把 Q 值拆成状态价值 V(s) 和动作优势 A(s, a) 之和。直觉是在很多状态下无论做什么动作结果差距不大这时候学习一个“状态本身有多好”比学习“每个动作有多好”更高效。class DuelingDQN(nn.Module): def __init__(self, input_channels4, num_actions6): super().__init__() self.conv nn.Sequential( nn.Conv2d(input_channels, 32, kernel_size8, stride4), nn.ReLU(), nn.Conv2d(32, 64, kernel_size4, stride2), nn.ReLU(), nn.Conv2d(64, 64, kernel_size3, stride1), nn.ReLU(), ) self.fc_value nn.Sequential( nn.Linear(64 * 7 * 7, 512), nn.ReLU(), nn.Linear(512, 1), ) self.fc_advantage nn.Sequential( nn.Linear(64 * 7 * 7, 512), nn.ReLU(), nn.Linear(512, num_actions), ) def forward(self, x): x x.float() / 255.0 features self.conv(x).reshape(x.size(0), -1) value self.fc_value(features) advantage self.fc_advantage(features) # 减去优势均值保证可辨识性 return value advantage - advantage.mean(dim1, keepdimTrue)最后一行是关键advantage - advantage.mean(dim1, keepdimTrue)是为了解决价值流和优势流之间的不可辨识性。如果不减均值同一个 Q 值可以对应无数种 V 和 A 的组合训练不稳定。优先经验回放PER解决的是样本利用率问题。普通回放池均匀采样但实际经验的价值并不相等——比如一次“丢球”的样本比一次“接住球”的样本更有信息量。PER 给每个样本分配一个优先级优先级越大越容易被采样。class PrioritizedReplayBuffer: def __init__(self, capacity, alpha0.6, beta_start0.4, beta_frames100000): self.buffer deque(maxlencapacity) self.priorities deque(maxlencapacity) self.alpha alpha self.beta beta_start self.beta_frames beta_frames def push(self, state, action, reward, next_state, done, td_error1.0): self.buffer.append((state, action, reward, next_state, done)) priority (abs(td_error) 1e-6) ** self.alpha self.priorities.append(priority) def sample(self, batch_size): probs np.array(self.priorities) probs probs / probs.sum() indices np.random.choice(len(self.buffer), batch_size, pprobs) batch [self.buffer[i] for i in indices] weights (np.array(self.priorities)[indices]) ** (-self.beta) weights weights / weights.max() return batch, indices, weights这个实现里alpha控制优先级的使用程度beta控制重要性采样权重的补偿程度。刚开始训练时beta较小比如 0.4逐渐升到 1.0这样既能利用优先采样又不至于让梯度产生太大偏差。实际使用时训练时需要有额外的 TD error 来更新优先级这就要在update_policy里额外计算一次td_error (targets - q_values).detach()再回传给 buffer 的update_priority方法。三个变体的组合方式非常灵活可以只用 Double也可以 Double Dueling再叠加 PER。业界常说的 Rainbow DQN 就是把六个改进全部合在一起但如果你想逐项验证效果从单个变体开始改更容易定位问题。4. 训练效果评估与可视化看曲线之外还要看真实回放4.1 评估协议每 N 步保存一次模型用确定性策略跑 100 局训练过程中光看训练回报曲线是不够的因为训练用的 epsilon 在变化回报会随 epsilon 波动而波动。正确评估方式是周期性暂停训练用epsilon0完全贪心跑若干局记录平均得分。def evaluate(env, policy_net, episodes100, renderFalse): rewards [] for _ in range(episodes): state, _ env.reset() total_reward 0 done False while not done: with torch.no_grad(): q_values policy_net(torch.tensor(state[None]).float()) action q_values.argmax(dim1).item() state, reward, terminated, truncated, _ env.step(action) total_reward reward done terminated or truncated rewards.append(total_reward) return np.mean(rewards), np.std(rewards)注意评估时环境不要用训练时的同一个实例因为训练环境已经被 wrapper 修改过比如 reward clip、episode life。最好单独创建一个干净的环境来做评估否则评估结果会受训练障碍影响。保存模型时除了保存权重还要保存训练超参数和评估结果torch.save({ model_state_dict: policy_net.state_dict(), optimizer_state_dict: optimizer.state_dict(), epsilon: epsilon, eval_reward: eval_reward, }, fcheckpoints/dqn_{timestep}.pt)这样断点续训时可以直接加载不必从头再来。这个习惯在长时间训练时非常重要——我见过不少人跑了两天发现曲线不涨然后发现是最开始学习率设错了但模型没存只能重跑。4.2 用 TensorBoard 记录训练指标Loss、Q 值、epsilon、每秒步数训练 DQN 时除了回报曲线我还要盯着四个指标TD loss、平均 Q 值、平均 reward、每秒步数FPS。TD loss 下降说明网络在收敛但如果 Q 值同时也在疯狂上升那可能是过高估计在作怪。每秒步数则是排查性能瓶颈用的如果只有个位数 FPS说明预处理或网络前向太慢。from torch.utils.tensorboard import SummaryWriter writer SummaryWriter(runs/dqn_pong) # 训练循环内 if total_steps % 1000 0: writer.add_scalar(train/loss, running_loss, total_steps) writer.add_scalar(train/q_value, avg_q_value, total_steps) writer.add_scalar(train/reward, episode_reward_avg, total_steps) writer.add_scalar(train/epsilon, epsilon, total_steps) # 评估后 writer.add_scalar(eval/mean_reward, eval_reward, total_steps)TensorBoard 是查看这些指标最省事的工具。不过要注意loss 的绝对值本身意义不大因为 targets 随网络变化而变化关键是看是否在某个区间内波动并逐渐变小。如果 loss 不断震荡上升那大概率是学习率太大了。4.3 保存游戏回放用 gymnasium 的 render_mode 记录智能体行为曲线和数据都是抽象的想真正判断智能体有没有学会策略最好直接看游戏画面。在评估时把render_modergb_array的环境分段保存成 mp4或者用gymnasium.wrappers.RecordVideo自动生成视频from gymnasium.wrappers import RecordVideo eval_env gym.make(ALE/Pong-v5, render_modergb_array) eval_env RecordVideo(eval_env, video_foldervideos/dqn_eval, episode_triggerlambda episode: episode % 10 0)RecordVideo会在触发条件满足时自动录制当前 episode 的视频存储在指定文件夹。这个功能特别适合用来诊断“看起来在动但实际在乱打”的模型。我通常保存三组视频训练刚开始探索期、训练中期、训练结束。对比看你会发现初期智能体在乱撞中期学会了回位后期会预判球的方向。如果中途发现视频里智能体老是朝一个固定方向按那就是动作空间或者网络输出有 bug比如argmax写错了轴。5. DQN 调参与避坑五个让训练失败的常见原因及修正方法5.1 现象loss 一直不降reward 一直在零附近徘徊原因之一是最常见的“奖励完全为 0”问题。Atari 环境里有些游戏只有得分变化才有奖励比如Breakout在没打中砖块前奖励恒为 0。如果 reward 的绝对值非常小loss 也会很小但这不代表网络在学东西因为所有 Q 值都趋近于 0。解决的方法有三个提高 reward scale把 reward 乘以 10或者使用经过life截断的 episode 来增加有效信号或者换一个奖励更密集的 game 来调试代码。我建议一开始先用Pong调试因为 Pong 几乎每几步就有 1 或 -1 的反馈训练节奏快方便验证代码正确性。5.2 现象训练早期 loss 飙高然后突然跌成 0模型不再输出动作这大概率是done标志处理错了。注意terminated和truncated是两回事terminated是游戏真的结束了truncated是到达最大步数被截断。如果只把terminated当成 done而忽略了truncated那么在截断时下一状态会被当成有效状态继续学习目标值将包含虚构的未来奖励。解决方法是像我在主循环里写的那样done terminated or truncated。还有一个低级错误在buffer.push时把done存成了 int然后训练时直接乘以(1 - dones_t)但 numpy bool 和 int 转换容易踩坑。统一用float32存储dones并且记得转换成 tensor 后再参与乘法运算。5.3 现象用state[None]时报维度错误或者训练步数极慢维度错误多出在预处理后状态形状不一致。比如AtariPreprocessor里reset返回的self.frames是 (4, 84, 84)而step里self.frames[1:]是 (3, 84, 84)拼上frame[None, ...]后变成 (4, 84, 84)形状没问题。但如果某个环境返回的观测是 (210, 160, 3)而你在_preprocess里用了cv2.COLOR_RGB2GRAY请确认观测确实是 RGB 顺序有时候 gymnasium 返回的是 BGR你需要看文档或用cv2.COLOR_BGR2GRAY来匹配。步数慢的排查路径先看 FPS如果 FPS 极低个位数检查是不是在训练循环里做了 GPU 同步比如每步都调用了torch.cuda.synchronize()。另一个隐蔽的坑是状态数组是np.uint8但每次训练时都torch.tensor(states).float()强制转成 float这会频繁分配内存。更好做法是预处理时直接用float32或者在缓冲区存储时就转 float免得每步都转换。5.4 现象训练很久reward 曲线上去了但评估时得分反而更低这是典型的“训练环境与评估环境不一致”问题。训练时你用了EpisodicLifeEnv、reward clipping、frameskip 等包装评估时如果仍用同一套包装评估得到的分数就会和训练分数混在一起失去独立评估的意义。正确做法是训练环境用一个包装好的实例评估环境单独创建一个只用必要预处理的实例。另外评估时epsilon必须为 0不然随机动作会拉低得分导致你以为模型没学会。5.5 现象Double DQN 反而比 DQN 效果更差这可能是因为网络容量不足或训练步数不够。Double DQN 在一定程度上降低了 Q 值的过估计但也让目标更新变得“更保守”初期学习速度会比 DQN 慢一点。这不是 bug是特性。如果你发现 Double DQN 在 20 万步内不如 DQN别急着回退让它跑满 50 万步再看。如果长时间依然落后检查是不是在线网络和目标网络同步太频繁target_update步数设置太短会导致目标变化太快建议设为 1000 到 10000 之间。另外有些游戏本身就没见过足够多的状态Double 的优势体现不出来。我一般在Pong、Breakout这类状态可分的游戏上验证变体效果而在BankHeist这类奖励稀疏的游戏上再用 Dueling 或 PER 来补足。6. 进阶把训练好的 DQN 从实验代码变成可评估的工程模块训练出模型只是第一步如何把模型沉淀成一个可复用、可对比、可断点续训的模块才是接近生产环境的做法。我常用的做法是把整个流程整理成一个config字典和三个函数create_env、create_agent、train然后用命令行参数覆盖配置。这样不仅能快速切换游戏还能把不同变体的超参数对比记录下来。一个实用的小技巧是用argparse定义所有超参数并把训练日志输出到带时间戳的目录import argparse parser argparse.ArgumentParser() parser.add_argument(--env, defaultALE/Pong-v5) parser.add_argument(--algo, defaultdqn, choices[dqn, double, dueling, rainbow]) parser.add_argument(--lr, typefloat, default1e-4) parser.add_argument(--batch_size, typeint, default32) parser.add_argument(--gamma, typefloat, default0.99) parser.add_argument(--replay_size, typeint, default100000) parser.add_argument(--target_update, typeint, default1000) parser.add_argument(--train_steps, typeint, default500000) args parser.parse_args()对于模型验证我拿手的是“对比同一环境、同一随机种子下不同变体的曲线”。每跑一个配置都把eval_reward记录到一个 CSV 里最后画成对比图。这个对比图比任何口头结论都有说服力。训练完成后把模型导出成 ONNX 格式以便在其他环境里做推理dummy_input torch.zeros(1, 4, 84, 84) torch.onnx.export(policy_net, dummy_input, dqn_pong.onnx, input_names[obs], output_names[q_values], dynamic_axes{obs: {0: batch}, q_values: {0: batch}})导出 ONNX 的价值在于你可以脱离 PyTorch 运行用 ONNX Runtime 在 CPU 上做推理速度比完整加载 PyTorch 模型快很多也方便嵌入其他系统。最后说一说我的习惯跑 DQN 实验我会把超参数、随机种子、环境版本、代码 git commit 哈希一起写进日志文件。这不是形式主义而是为了防止过两周你看到一条不涨的曲线却忘了当时用的哪个版本和哪份代码。强化学习的“玄学”往往就藏在版本差异和随机种子差异里。整个方案下来我最深的体会是DQN 及其变体在 Atari 上实现并不难难的是建立一套可复现、可对比的实验流程。当你把环境封装、回放池、目标网络、评估协议都一次写对训练曲线自然会给你正向反馈。希望这篇落地笔记能让你少走几个弯路早点跑到那条朝上的曲线。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联
返回资讯列表 →