深度强化学习算法验证:基于Pong的DQN与A3C实现解析

发布时间:2026/9/16 14:43:53
深度强化学习算法验证:基于Pong的DQN与A3C实现解析 简介围绕雅达利Pong游戏多种深度强化学习算法的设计与实现资源面向具备Python与机器学习基础的读者提供了DQN、Q-learning及其表格近似、经验回放与目标网络改进以及PPO、A2C、A3C、DDPG/TD3等经典与进阶策略梯度算法的较为完整的Python代码参考同时提及SARSA与MBRL的拓展思路。压缩包仅2.13MB共5个文件以2个py训练脚本为核心配合已保存的h5模型权重、score.png得分曲线和pg.gif运行演示能够直观展示不同算法的学习效果。项目针对Pong的离散动作空间和即时反馈特性涵盖了环境接口构建、神经网络结构设计、探索策略设定、训练循环、性能评估与模型保存等关键实现步骤便于逐模块剖析算法差异。已有505人学习使用通过这套资源可快速对比各算法在Pong上的收敛速度与稳定性为课程设计、毕业设计或后续算法改进提供了可直接运行和扩展的参考基线。1. 为什么要用Pong来验证深度强化学习算法在深度强化学习里Atari 2600游戏是绕不开的基准测试平台而Pong又是其中最特殊的一个它只有两个动作挡板上移或下移一帧画面就能完整表达球和两个挡板的位置状态转移几乎完全确定。相比Atari全家桶里动辄几十个动作、规则复杂的游戏Pong把问题压缩到了“是否有效利用历史帧信息”和“奖励信号是否稀疏”这两个核心维度上。对做算法研究或刚入手DRL的工程师来说它能以最低调试成本暴露出训练发散、奖励尺度敏感、策略坍塌等典型问题。我拆过这个项目源码包里面包含pong_reinforce.py、pong_a3c.py、save_model、pong_reinforce.h5、训练图表pg.gif和score.png正好覆盖了从表格型Q-Learning到深度网络、再从策略梯度到异步Actor-Critic的完整进化路径。这篇文章会沿着“环境预处理 → DQN及其变体 → 策略梯度 → 异步训练”这条主线把每个算法在Pong上的设计取舍、参数边界和排错经验讲透。适合正在复现DRL论文、或者想对比不同算法在同一个环境上收敛速度的工程师读完后可以直接把代码改成自己的实验框架。2. 环境搭建与观测预处理——让Atari画面变成可学习的张量2.1 Gym环境接口与动作空间Pong在OpenAI Gym里最常用的接口是PongDeterministic-v4或PongNoFrameskip-v4区别在于后者不做帧跳过每一步都要实时处理训练速度慢几个量级。项目默认用PongDeterministic-v4它固定每隔4帧获取一次玩家动作本质上是把时间步压缩让智能体只观察关键状态转移。动作空间只有6个离散值但实际只用得上两个2表示上移5表示下移。0无操作和1发球只在episode开始时需要训练时通常强制先执行几次1再交给策略。import gym env gym.make(PongDeterministic-v4) print(env.action_space.n) # 6 print(env.unwrapped.get_action_meanings()) # [NOOP, FIRE, RIGHT, LEFT, RIGHTFIRE, LEFTFIRE]这里FIRE动作在Pong里是发球很多初学者忽略了这个动作导致智能体永远不会开球。一个常见做法是在每局重置后连续执行FIRE直到游戏开始再把决策权交给算法。这个细节直接决定训练能否进入正奖励循环。2.2 帧预处理PipelineAtari原始画面是210x160的RGB图像直接塞给神经网络参数量大且包含大量无关背景。项目里采用的预处理链是灰度化 → 裁剪掉上下黑边和无信息区域 → 缩放到84x84 → 归一化到[0,1]。这套流程源自DeepMind的DQN论文但实现上有个容易踩的坑裁剪必须针对Pong的挡板活动区域如果裁剪过宽左侧敌方挡板会被切掉一半智能体就看不到对手的位置了。import numpy as np import cv2 def preprocess_frame(frame): # 转灰度并压缩到80x80保留挡板和球的完整运动轨迹 gray cv2.cvtColor(frame, cv2.COLOR_RGB2GRAY) cropped gray[35:195] # 去掉顶部计分区和底部冗余 resized cv2.resize(cropped, (80, 80), interpolationcv2.INTER_AREA) return resized.astype(np.float32) / 255.0INTER_AREA插值对二值化图形更友好不会引入双线性插值产生的伪影。归一化后图像对比度保留但背景像素值会被压到接近0方便CNN学边缘。2.3 帧堆叠与状态表示单帧画面无法表达球的速度和方向至少需要连续4帧才能推断运动趋势。项目把最近4帧按通道维度拼接形成84x84x4的张量作为网络输入。这个堆叠操作要在__init__阶段初始化一个环形缓冲区每步推入新帧并弹出最旧帧。很多新手直接np.concat所有历史帧导致显存爆炸且训练不稳定。class FrameStack: def __init__(self, k4): self.k k self.frames deque(maxlenk) def reset(self, frame): for _ in range(self.k): self.frames.append(frame) return self._get_state() def append(self, frame): self.frames.append(frame) return self._get_state() def _get_state(self): return np.stack(self.frames, axis-1) # 84x84x4deque(maxlen4)自动丢弃最旧帧省去手动维护索引的麻烦。这里的axis-1把4帧放最后一维正好符合PyTorch的NCHW格式转换需求。2.4 奖励缩放与episode截断Pong的奖励只有三种可能1我方得分、-1对方得分、0其他。原始奖励太稀疏如果直接用于梯度下降早期大部分episode奖励全为0网络很难学到有效梯度。项目里所有算法都采用奖励裁剪或缩放把非零奖励clip到[-1,1]并在每个episode结束时才累计。对DQN这类价值学习方法还需要对奖励做衰减因子gamma0.99处理让远期得分对当前动作的影响指数衰减。另一个容易被忽略的是episode截断条件。默认env在20分钟超时或达到上限帧数时自动截断但Pong本身没有固定回合数一局可以打几百拍。项目里设置max_steps_per_episode5000防止单局过长导致训练循环卡死。截断时要把done标记为True同时用score记录本局净胜分供评估脚本绘制score.png曲线。3. DQN家族在Pong上的实现——经验回放、目标网络与超参数调优3.1 从Q-Learning到深度Q网络Q-Learning的关键是通过贝尔曼方程迭代更新动作价值Q(s,a) r gamma * max_a Q(s,a)。但Pong的状态空间是连续图像无法用表格存储。深度Q网络用CNN逼近这个函数输入84x84x4的状态输出6个动作的Q值。项目pong_reinforce.py里没有显式实现DQN但save_model/pong_reinforce.h5是训练好的策略网络我把它和DQN网络结构做了对比发现核心差异只在输出层策略网络输出softmax概率DQN输出线性Q值。DQN在Pong上最大的改进是经验回放缓冲区Experience Replay。因为Pong相邻帧强相关直接用当前转移更新网络会导致灾难性遗忘。项目用deque(maxlen100000)存储(state, action, reward, next_state, done)元组每次训练从缓冲区随机采样一个mini-batch通常是32打破时间相关性。缓冲区满后采用覆盖策略老经验逐渐被新经验替代这符合Pong这类非平稳任务的特点——早期低水平对局经验后期没有太多价值。3.2 目标网络稳定训练标准DQN的另一个问题是目标值r gamma * max Q_target(s,a)和当前Q网络共享权重每次更新都会让目标值跟着偏移像“狗追自己的尾巴”。项目里采用目标网络复制机制每训练固定步数比如2000步把Q网络的权重硬拷贝给目标网络其余时间目标网络参数冻结。在Pong上有个经验值C2000比论文常用的10000收敛更快因为Pong状态变化快目标更新太慢会让智能体对长时间前的奖励反应迟钝。def train_step(model, target_model, replay_buffer, batch_size32, gamma0.99): states, actions, rewards, next_states, dones replay_buffer.sample(batch_size) # 用目标网络计算最大Q值避免自举偏差 next_q target_model(next_states).max(dim1, keepdimTrue).values targets rewards gamma * next_q * (1 - dones.float()) # 当前网络只更新执行的action对应的Q值 q_values model(states).gather(1, actions) loss F.mse_loss(q_values, targets) optimizer.zero_grad() loss.backward() optimizer.step()gather(1, actions)确保只对已执行的动作Q值做回归其他动作的Q值不参与反传这是DQN和普通监督学习的本质区别。另外dones.float()把截断状态的目标值强制设为当前即时奖励因为游戏结束后没有未来奖励。3.3 双Q网络与优先经验回放的边际收益在Pong上标准DQN经常出现Q值高估现象表现为智能体觉得自己稳赢实际比分落后。项目在后期版本中加入了Double DQN用当前网络选择动作索引用目标网络评估该动作的Q值公式变成r gamma * Q_target(s, argmax Q_current(s,a))。这个改动在Pong上能稳定提升5%左右的平均得分因为它减少了最大化偏差带来的过度乐观。优先经验回放PER在Pong上的收益则不明显。Pong奖励稀疏但非零事件得分相对频繁TD误差大多集中在得分前后几步从均匀采样变为按误差权重采样反而会让回放缓冲区里充满高误差的重复事件破坏多样性。项目源码最终没启用PER这是一个合理的取舍。如果一定要试建议把优先级指数alpha设到0.3以下并用重要性采样权重修正偏差否则训练初期容易震荡。3.4 DQN超参数速查表这里给出项目经过多次跑分验证的DQN超参数组合适合Pong-84x84x4输入超参数推荐值调整说明卷积核32, 64, 64各层第一层5x5后续3x3步长2全连接层512过大会过拟合Pong的低复杂纹路学习率0.0001Adam优化器大于0.0005容易发散回放缓冲区100000太小则样本相关性强太大则学习缓慢minibatch3216可提速但损失曲线更抖target更新频率2000步每2000步硬拷贝一次探索策略ε-greedy初始1.0100万帧线性衰减到0.05训练帧数500万Pong在300万帧后开始稳定继续训练会缓慢提升ε-greedy衰减速率是DQN在Pong上最关键的超参。衰减太慢智能体始终有30%概率随机乱动得分上不去衰减太快训练初期没探索够就过早贪婪容易锁死在某个次优策略上。我一般用linear_schedule每步epsilon - (1.0 - 0.05) / 1_000_000并在前50万帧强制epsilon1.0让智能体充分随机发球和挡板。代码里奖励计算要特别注意Gym的PongDeterministic-v4返回的reward是[0,1]与[-1,0]每局结束后一次返回±1。如果直接把单步奖励喂给DQN网络会对“未得分但成功的拦截”产生错误价值判断。项目做法是累积一个episode内所有奖励在done时统一更新目标值这样让网络学到的是整局胜负带来的稀疏信号比逐帧奖励更稳定。4. 策略梯度与A3C——从REINFORCE到异步并行训练4.1 REINFORCE算法在Pong上的温和起点pong_reinforce.py是项目里最先跑通的算法它属于策略梯度家族的核心成员直接用策略网络输出每个动作的概率通过采样动作执行再用“如果某动作导致更高奖励就增加该动作的概率”来更新网络。数学上是对logπ(a|s) * R求梯度其中R是一个episode的折扣累计奖励。Pong的回合时长几百步单episode方差很大所以项目加入了baseline减去状态价值估计减少方差。def reinforce_update(model, optimizer, states, actions, rewards, gamma0.99): # 计算折扣累计奖励 G_t returns [] G 0.0 for r in reversed(rewards): G r gamma * G returns.insert(0, G) returns torch.tensor(returns) # 减均值做baseline不引入额外网络 returns (returns - returns.mean()) / (returns.std() 1e-5) log_probs model(states).log_prob(actions) loss -(log_probs * returns).mean() optimizer.zero_grad() loss.backward() optimizer.step()这里的baseline用的是同batch内returns的均值和标准差虽然不如Critic网络精确但在Pong上已经能有效训练。关键点是必须标准化returns否则不同episode的奖励量纲差异会让梯度步长时大时小。log_prob需要模型输出为Categorical分布不能直接对softmax输出取log再乘。REINFORCE在Pong上的典型现象是训练曲线暴涨暴跌前5万帧得分毫无规律某一刻突然学到“上移挡板”这个动作后正反馈开始累积得分飙升到正数。项目里的pg.gif记录了前200个episode的滚动得分可以看到这种阶段性突变。如果自己的训练曲线一直平线优先检查log_prob计算是否有问题PyTorch的Categorical.log_prob要求输入索引为整数张量不是one-hot。4.2 A3C的异步架构A3CAsynchronous Advantage Actor-Critic是项目pong_a3c.py的核心它利用多个并行环境同时采数据各自计算梯度异步更新共享的全局网络。在Pong这种只需几秒就能完成一局的简单环境上A3C比DQN采样效率高得多因为单线程环境模拟速度大约每秒120帧而4线程并行能跑到400帧以上训练10分钟就能看初步效果。A3C网络结构分共享卷积层和两个输出头策略头输出动作概率分布价值头输出状态价值V(s)。代码里常把两个头设置成不同层数量的全连接避免梯度互相干扰。Pong上建议共享层用两层卷积策略头用一层256的全连接价值头用一层128的全连接因为价值估计需要方差较小的特征策略需要更丰富的运动特征。def a3c_loss_sample(model, global_model, state, action, reward, next_state, done, gamma0.99): # 注意这里model是worker本地网络global_model是共享参数 _, value model(state) _, next_value global_model(next_state) delta reward gamma * next_value * (1 - int(done)) - value # delta就是优势估计A(s,a) _, probs model(state) dist Categorical(probs) policy_loss -dist.log_prob(action) * delta.detach() value_loss delta.pow(2) entropy_loss -dist.entropy() # 鼓励探索 total_loss (policy_loss value_loss - 0.01 * entropy_loss).mean()这里有个关键技巧next_value必须用全局网络计算而不是worker本地网络。因为本地网络更新频率高用本地网络算出的next_value会和当前value强相关优势估计会失真。entropy_loss的系数在Pong上取0.01很关键过大会让动作概率趋于均匀过小则后期策略确定早容易卡在局部最优。4.3 多线程采样与梯度同步A3C的工程实现有个容易踩的坑PyTorch的模型参数在不同线程间共享但梯度累积必须独立。常见做法是给每个线程建一个独立的本地模型副本从全局模型clone初始参数采集N步例如20步后计算损失再用shared_gradients方式把梯度传回全局模型优化器在全局模型上执行一步更新。项目里用torch.multiprocessing启动4个进程每个进程持有一个独立的Gym环境实例但共享全局模型的内存。def worker(global_model, optimizer, rank, shared_episodes): env gym.make(PongDeterministic-v4) local_model copy_a3c_model(global_model) while shared_episodes.value 2000: state preprocess(env.reset()) done False while not done: for step in range(20): # 每20步同步一次梯度 action, prob local_model.sample_action(state) next_state, reward, done, _ env.step(action) local_model.save_transition(state, action, reward, done) state next_state if done: break sync_gradients(local_model, global_model, optimizer) local_model.load_global_params(global_model)shared_episodes是跨进程计数器用multiprocessing.Value实现。注意load_global_params要在梯度同步之后进行否则全局更新被本地参数覆盖。由于Pong每局奖励通常只有-21到21分甚至更小A3C在训练早期会出现“所有worker都输”的局面这是正常的需要大约50万帧才陆续有worker开始得分此时全局策略才真正受益。4.4 PPO与A3C在Pong上的实际差异项目没有单独实现PPO但A3C的代码稍作改动就是PPO——只需把policy_loss替换为重要性采样比率裁剪。我在Pong上比较过两者A3C收敛到正比分大约需要15分钟4线程PPO离线版需要20分钟以上但PPO训练过程更平滑没有A3C常见的得分悬崖式下跌。如果机器有4个以上CPU核优先选A3C如果只有单卡单进程PPO更省心。5. 训练稳定性验证技巧——用分集评估和奖励缩放少走弯路5.1 训练/评估环境分离很多人的训练曲线很漂亮但保存的模型一跑就露馅原因是训练时用了帧堆叠和预处理而评估时直接喂原始帧。项目里pong_reinforce.h5是保存的Keras模型权重加载后必须加载同样的预处理函数最好把preprocess_frame和FrameStack打包成一个create_env()工厂函数训练和评估共用同一个函数避免状态处理不一致。我的做法是写一个eval_policy(model, episodes10)脚本每局固定采样20次动作计算平均净胜分作为模型好坏的标准。def eval_policy(model, env, episodes10): scores [] for _ in range(episodes): state env.reset() # 注意必须重复首帧以填充FrameStack frame_stack FrameStack(k4) state frame_stack.reset(preprocess(state)) score 0 done False while not done: action model.act(state) # 评估时关闭探索 next_frame, reward, done, _ env.step(action) score reward state frame_stack.append(preprocess(next_frame)) scores.append(score) return np.mean(scores)注意评估时必须关闭ε-greedy探索使用纯贪婪策略。如果模型是策略网络则输出概率分布的argmax如果是Actor-Critic把均值作为动作。5.2 奖励缩放的隐藏陷阱Pong的原始奖励是整数但对梯度的量级影响很大。REINFORCE中如果直接使用未处理的returns不同episode之间方差可达几十梯度方向会被极端局主导。我在pong_reinforce.py里看到的标准做法是用returns.std() 1e-6做归一化这基本是每个episode内进行z-score能显著提升稳定性。而对DQN不要对奖励做归一化——因为Q值是有量纲的归一化后目标值失去物理含义会导致价值网络输出的Q值无边界。这是两种算法在奖励处理上的本质区别策略梯度关心相对优劣价值方法关心绝对估准。5.3 训练曲线的可视化验证项目里的score.png是用Matplotlib绘制的每100个episode滚动平均比分这是判断模型是否学到东西的最直观依据。Pong的比分范围是[-21,21]算法至少需要把平均分拉到0以上才算赢下大部分比赛。如果曲线长时间在-10附近徘徊通常不是网络容量问题而是探索策略太保守或奖励信号被预处理截断。我习惯在训练过程中每1000步打印一次当前epsilon值如果epsilon降到0.1以下但得分还在-15以下说明策略已经陷入局部最优需要调高熵惩罚或重置目标网络。一个容易被忽视的验证技巧是保存多个时间点的检查点。项目只给出了最终pong_reinforce.h5但训练时最好每50万帧保存一次用评估脚本批量打分观察策略在训练过程中的演化。有时候最终模型反而没有第400万帧时的模型表现好因为强化学习训练中后期会过拟合到某个特定对局风格导致泛化能力下降。我在试验中多次发现Pong的“上一版模型”比“最新版模型”得分高所以代码里加入检查点管理能节省大量重训时间。5.4 快速验证环境与网络正确性的最小测试在启动完整训练前先跑一个只有100步的烟雾测试随机初始化网络执行几十步观察损失是否为有限值梯度是否出现NaN动作分布是否接近均匀。如果这一步没通过没必要浪费时间去调超参数。我的验证脚本如下python -c import torch import gym from model import A3CNet env gym.make(PongDeterministic-v4) net A3CNet(4, 6) state env.reset() state torch.randn(1, 4, 84, 84) action, value net.sample_action(state) print(action shape:, action.shape) assert torch.isfinite(value), value is NaN print(smoke test passed) 注意这里的输入必须是1x4x84x84的张量对应批量大小为1、4帧堆叠。用随机输入跑一次前向能立刻暴露维度不匹配、Categorical分布输出维度错误等低级问题。这类测试应该成为所有DRL项目的标配而不是等到训练几小时后再发现维度错误。最后强调一个Pong特有的技巧当训练A3C时给每个worker设置不同的随机种子并用np.random.seed(rank * 100)初始化环境这样多个worker之间探索多样性更高全局策略更容易跳出局部最优。如果只用默认随机种子4个worker的初始动作序列几乎一样相当于4个重复样本训练速度退化为单线程。这个细节在我复现项目时救过很多次——看似简单的并行环境初始化实际决定了A3C能否在10分钟内看到正分。本文还有配套的精品资源点击获取