尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

DQN深度强化学习算法详解:从Q-Learning到PyTorch实战

DQN深度强化学习算法详解:从Q-Learning到PyTorch实战 DQN算法算是我在深度强化学习这条路上真正跑通的第一个算法也是很多新手入门强化学习时绕不开的一道坎。作为一个曾经被Q值发散、reward震荡折磨得睡不着觉的人我准备把DQN从原理到代码实现整个过程掰开了讲一遍希望能帮你少踩几个坑。这篇文章适合有Python和PyTorch基础、想系统学习深度强化学习或者正在复现论文却卡在DQN细节上的同学看完之后你不仅能亲手实现一个能跑的DQN还能理解它背后的每一个设计逻辑。1. 为什么要理解DQN从Q-Learning到深度网络的进化逻辑1.1 DQN要解决的到底是什么问题要理解DQN得先知道它的前身Q-Learning。传统Q-Learning的核心是用一张Q表存储状态-动作价值表每次决策时查表选出Q值最大的动作。但这种方案的瓶颈非常致命状态空间一旦变大Q表会膨胀到无法存储更别说处理图像这类高维输入了。DQN的思路就很直接——用一个深度神经网络来代替这张Q表输入是状态比如游戏画面的原始像素输出是每个动作的Q值估计。换句话说DQN把“查表”变成了“函数拟合”让模型可以通过训练泛化到从未见过的状态。这是深度强化学习真正意义上打开局面的代表作。在动手写代码之前我建议你先想清楚一个问题DQN到底在训练什么它的训练目标不是让网络输出精确的Q值而是让Q值估计逐步逼近真实收益的期望。整个DQN的训练流程可以概括为用当前网络选择动作获得经验从历史经验中随机采样然后计算目标Q值并更新网络参数。这四个环节环环相扣任何一个环节出问题都会导致训练失败。1.2 让DQN真正稳定起飞的三个关键机制第一次实现DQN的人最直观的感受就是为什么我的网络训练着训练着就崩了原因在于强化学习的数据不是独立同分布的。相邻时间步的状态高度相关而当前的更新又会改变后续样本的分布这就是所谓的样本相关性和非平稳问题。DQN引入了三个关键机制来应对上述问题经验回放Experience Replay把智能体与环境交互的样本存进一个固定大小的回放缓冲区训练时从中随机采样一个小批量而不是使用相邻时间步的数据。这个操作的巧妙之处在于随机采样打乱了样本间的时序相关性同时一条经验还能被反复利用提升了样本效率。目标网络Target Network训练时需要计算目标值 ( r \gamma \max_{a} Q(s, a) )如果每次都直接用当前正在更新的网络计算最大值目标会随着网络参数的变化而不断移动就像追着自己的尾巴跑很难收敛。DQN的做法是复制一份参数冻结的目标网络每隔固定步数才同步一次让训练目标在一段时间内保持稳定。误差截断把损失函数中的误差项限制在 ([-1, 1]) 范围内当TD误差过大的时候把梯度限制住防止参数出现剧烈抖动。这三个机制并非DQN论文里拍脑袋想出来的而是针对强化学习训练不稳定的三大病根逐一做的对症下药。理解这一点之后你后面调参时就不会盲目乱试了。2. 核心细节解析损失函数、网络结构与超参选择2.1 目标Q值的构建与损失函数推导DQN的损失函数本质上是最小化当前Q值与目标Q值的均方误差但它和普通的监督学习有个本质区别——标签不是事先给定的而是通过贝尔曼方程自举估计出来的。具体来说对于从经验池中采样得到的一批样本 ( (s, a, r, s, done) )目标Q值 ( y ) 按如下规则计算如果 ( s ) 是终止状态那么 ( y r )因为终止状态之后没有未来收益。否则 ( y r \gamma \max_{a} Q_ {\text{target}} (s, a) )其中 ( Q_ {\text{target}} ) 是目标网络( \gamma ) 是折扣因子。然后计算当前网络的预测Q值 ( Q_ {\text{online}}(s, a) )注意这里只取当前状态下实际执行动作对应的Q值。损失函数定义为两者的均方误差。在PyTorch里计算目标Q值的一个常见写法是这样的next_q_values target_net(next_states).max(dim1, keepdimTrue)[0] target_q_values rewards (1 - dones) * gamma * next_q_values这里有三个细节很容易踩坑dones需要与rewards形状一致并且要把它转成浮点数而不是直接用布尔值参与乘法运算。如果next_states有值但dones为True的样本被错误地参与了目标计算会导致Q值在终止状态处被高估。max(dim1)取的是每个状态在所有动作上的最大Q值这一步就是公式中的 ( \max_{a} ) 。在Double DQN中这个写法要改掉后文我会专门讲。2.2 网络结构设计与输入输出对齐DQN的网络结构并没有一个统一的标准它的设计完全取决于你面对的任务类型。以经典的控制任务CartPole为例输入是4维的状态向量位置、速度、角度、角速度输出是2个动作的Q值所以网络可以是简单的三层全连接网络。而如果你要做Atari游戏输入变成了84x84的灰度图像这时候就需要用卷积层来提取视觉特征然后接全连接层输出动作维度的Q值。我个人的经验是先根据输入类型定网络骨架再根据动作空间大小定输出维度两者对齐比网络深度更重要。有一点特别值得提醒输出层的激活函数千万不要用ReLU或者sigmoid。Q值的物理意义是长期累计收益的期望它本身就是一个无界实数所以输出层应该保持线性激活。一旦你给输出层加了限制Q值的取值范围就会被卡死模型永远无法学到正确的价值估计。2.3 关键超参数的经验取值与调整方向超参数在DQN里的重要性不亚于算法本身。我整理了一张基于CartPole环境的起点配置表这也是我多次运行后觉得比较稳妥的一组参数超参数推荐取值说明学习率1e-3过高会导致Q值发散过低收敛太慢折扣因子 ( \gamma )0.99控制模型对远期收益的重视程度经验池容量10000太小则样本多样性不足太大会拖慢采样批量大小32兼顾更新频率与稳定性目标网络更新间隔100步更新过快起不到稳定作用探索率 ( \epsilon ) 初始值1.0初始阶段以随机探索为主探索率衰减终点0.01随着训练推进逐步减少探索探索率衰减步数1000步具体衰减速度视任务复杂度调整epsilon的衰减策略是新手最常忽略的一项。我习惯用线性衰减从epsilon_start 1.0开始每步衰减(epsilon_start - epsilon_end) / epsilon_decay_steps衰减到epsilon_end 0.01后保持不变。这里有一个判断标准如果你发现智能体在训练后期仍然大量随机探索说明epsilon衰减得太慢了反过来如果一上来就变成纯利用模式很容易陷入局部最优。2.4 奖励设计对训练效果的影响很多人会觉得DQN对奖励的容忍度很高只要环境给出正负信号就能学习。但实际上奖励的尺度和密度直接决定了训练的难度。CartPole环境自带的奖励是每坚持一步得1分这种持续稳定的正反馈让DQN比较容易学到保持平衡的策略。如果面对的是一个稀疏奖励环境比如只有到达终点才给一个1其他时刻都是0DQN就会面临非常严重的信用分配问题。一个实用的做法是给中间过渡状态加一个小的形状奖励reward shaping比如用距离目标点的变化量作为辅助信号帮助模型建立渐进的优化梯度。还有一个容易被忽略的点是奖励的绝对值大小。如果你的奖励动辄几百上千而网络的学习率还是1e-3梯度更新时很容易把参数冲到极差的位置。我通常会先把奖励做归一化处理让它们保持在合理的量级再喂给算法。3. 实战用PyTorch手写一个CartPole DQN3.1 环境准备与依赖说明动手之前先把环境装好。需要安装PyTorch和Gymnasium新版OpenAI Gym的维护分支API更友好推荐直接用它其他的就是常规的NumPy和Matplotlib。pip install torch gymnasium matplotlib numpy我用的是CartPole-v1环境它是一个连续控制任务智能体需要左右移动小车让杆子保持直立。状态空间是4维动作空间是2维每一步保持不倒得1分超过475分就算通关。这个环境足够简单非常适合验证DQN的实现是否正确。3.2 网络定义与经验池实现首先定义一个三层全连接网络输入层接收4维状态中间两个隐藏层各128个神经元激活函数用ReLU输出层是两个动作的Q值。import torch import torch.nn as nn import torch.nn.functional as F class DQN(nn.Module): def __init__(self, state_dim, action_dim, hidden_dim128): super(DQN, self).__init__() self.fc1 nn.Linear(state_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, hidden_dim) self.fc3 nn.Linear(hidden_dim, action_dim) def forward(self, x): x F.relu(self.fc1(x)) x F.relu(self.fc2(x)) return self.fc3(x)经验池可以基于collections.deque实现它天然支持固定长度队列超出长度后最早的数据会被自动弹出。我习惯再封装一个sample方法从队列中随机取出一个批量的样本并转换成PyTorch张量。from collections import deque import random import numpy as np class ReplayBuffer: def __init__(self, capacity10000): 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) state, action, reward, next_state, done map(np.stack, zip(*batch)) return ( torch.FloatTensor(state), torch.LongTensor(action).unsqueeze(1), torch.FloatTensor(reward).unsqueeze(1), torch.FloatTensor(next_state), torch.FloatTensor(done).unsqueeze(1), ) def __len__(self): return len(self.buffer)这里一个小细节action存储在buffer里时用的是整数类型但采样出来之后要转成torch.LongTensor因为后面要用它做索引来取特定动作的Q值。done必须转成浮点张量参与乘法运算不能停留在布尔类型。3.3 训练主循环与关键步骤拆解训练主循环是整个DQN实现的核心每一轮迭代都包含四个阶段环境交互、经验存储、参数更新、目标网络同步。我把它拆开来说。首先是从环境中采样状态根据epsilon-greedy策略选择动作。这个策略的表述非常简单以epsilon的概率随机选择一个动作探索否则通过当前Q网络选择Q值最大的那个动作利用。if random.random() epsilon: action env.action_space.sample() else: with torch.no_grad(): q_values online_net(torch.FloatTensor(state).unsqueeze(0)) action q_values.argmax().item()其次把交互得到的(state, action, reward, next_state, done)存入经验池。然后当经验池的数据量达到一定程度后从其中采样一个小批量用于训练。在写代码时需要注意如果经验池太小就开始训练样本的多样性不足模型很容易被少数极端样本带偏。我通常设定一个最低阈值比如1000条经验后才开始学习。参数更新阶段的代码写起来很规整核心逻辑如下batch replay_buffer.sample(batch_size) state, action, reward, next_state, done batch with torch.no_grad(): next_q target_net(next_state).max(dim1, keepdimTrue)[0] target_q reward gamma * next_q * (1 - done) current_q online_net(state).gather(1, action) loss F.mse_loss(current_q, target_q) optimizer.zero_grad() loss.backward() optimizer.step()这段代码里最重要的是gather函数的用法。online_net(state)输出的是一个batch_size x 2的张量而我们只关心实际执行的那个动作对应的Q值。gather(dim1, indexaction)就是从第二维根据index精确取出对应的值得到batch_size x 1的当前Q值估计。很多人第一次写的时候会用online_net(state)[action]这种方式但因为action是一个二维张量这样索引会非常容易出错。最后是目标网络同步。我习惯每个一定步数比如100步把目标网络的参数直接复制成在线网络的参数if step % target_update 0: target_net.load_state_dict(online_net.state_dict())3.4 完整训练流程与效果评估把上面的模块拼装起来一个可以完整跑通的训练脚本大致是这样的。我先初始化环境、网络、优化器和经验池然后进入循环累计每个episode的回报。为了观察训练效果我会每50个episode打印一次平均奖励并把训练过程中的reward曲线存下来。import gymnasium as gym env gym.make(CartPole-v1) state_dim env.observation_space.shape[0] action_dim env.action_space.n online_net DQN(state_dim, action_dim) target_net DQN(state_dim, action_dim) target_net.load_state_dict(online_net.state_dict()) optimizer torch.optim.Adam(online_net.parameters(), lr1e-3) replay_buffer ReplayBuffer(capacity10000) gamma 0.99 epsilon 1.0 epsilon_end 0.01 epsilon_decay 1000 batch_size 32 target_update 100 min_buffer_size 1000 episode_rewards [] step_count 0 for episode in range(500): state, _ env.reset() episode_reward 0 done False while not done: epsilon max(epsilon_end, epsilon - (1.0 - epsilon_end) / epsilon_decay) action ... # epsilon-greedy next_state, reward, terminated, truncated, _ env.step(action) done terminated or truncated replay_buffer.push(state, action, reward, next_state, done) state next_state episode_reward reward step_count 1 if len(replay_buffer) min_buffer_size: loss update() if step_count % target_update 0: target_net.load_state_dict(online_net.state_dict()) if done: break episode_rewards.append(episode_reward)在CartPole-v1上按这组参数配置通常在100-200个episode之内就能看到reward稳定超过300300个episode左右多数情况下能达到475以上的通关标准。如果你的结果迟迟上不去优先检查有没有犯第4节里提到的那些初始化错误而不是怀疑算法本身。4. 常见问题与排查技巧实录4.1 训练不收敛或者Reward一直很低这是DQN新手最容易碰到的问题。如果你发现跑了几百个episodereward曲线一直贴着地面先去检查这几处会不会是经验池里的样本从来没更新过模型参数需要确保if len(replay_buffer) min_buffer_size这个条件真的被触发了。CartPole一个episode大约能产生几到上百条经验如果你设的min_buffer_size太大很可能前几十个episode都在纯收集数据看起来就好像网络根本没在学。done标志有没有正确参与目标值计算这个bug非常隐蔽。终止状态下不应该有未来收益如果忘了对next_q乘以(1-done)终止状态附近的Q值会被错误地拉高模型会认为结束游戏是好事训练永远无法收敛。学习率是否过大DQN对学习率的敏感度比普通监督学习要高得多。我在一些任务上把学习率从1e-3调到1e-2Q值直接发散成NaN。遇到不收敛把学习率降一个数量级永远是最快验证手段。4.2 Q值过估计与Double DQN优化标准DQN使用了 ( \max ) 算子来估计目标Q值但这个max操作会带来系统性偏差——由于函数拟合误差的存在对Q值的估计往往偏高尤其在不确定性较大的区域。这种过估计问题overestimation在DQN中会导致智能体对某些动作产生盲目自信影响策略质量。解决思路其实不复杂Double DQN的核心是让动作选择和价值估计解耦。在计算目标Q值的时候用在线网络选动作但用目标网络去估计这个动作的Q值。也就是说不再直接取target_net(next_state).max()而是先online_net(next_state).argmax()再拿这个索引去target_net(next_state)里取值。with torch.no_grad(): next_actions online_net(next_state).argmax(dim1, keepdimTrue) next_q target_net(next_state).gather(1, next_actions) target_q reward gamma * next_q * (1 - done)这个改动就一行带来的稳定性提升却很可观。我在Atari环境上实测过Double DQN在多数游戏上的最终得分都优于标准DQN而且训练曲线更平滑。如果你的训练后期出现了reward骤降或者震荡加剧试着改成Double DQN。4.3 目标网络更新频率与经验池容量如何平衡目标网络更新频率和经验池容量之间实际存在一个微妙的平衡。更新太频繁会让训练目标不断移动起不到稳定作用更新太慢则目标网络和在线网络的Q值差距过大初期训练效率会受影响。我观察到的经验值是更新间隔大约在100-500步之间比较合适具体取决于你的环境交互频率和数据量。经验池容量也是类似逻辑。容量太小采样的样本集中在最近的经验相关性高起不到回放打破相关性的作用容量太大很多旧样本对应的策略已经过时也可能干扰学习。CartPole这种简单任务用10000就够了但如果是Atari这类复杂任务通常需要100000甚至更大的容量。一个实用的判断方法是看训练曲线的平滑度如果曲线抖动剧烈优先降低学习率或者增大经验池容量如果曲线一段时间完全不动再考虑是不是目标网络更新太慢导致梯度方向长期失真。4.4 从CartPole迁移到Atari类任务时的适配要点很多人在CartPole上跑通了DQN信心满满地迁移到Atari游戏上结果发现完全训不动这是很正常的。CartPole的状态是低维数值向量奖励稠密动作效果立竿见影而Atari的输入是高维图像奖励稀疏时间延迟长两者完全不在一个难度等级。从CartPole迁移到Atari你需要改动的地方不只是网络结构。首先输入图像要预处理成84x84的灰度图并且需要堆叠最近4帧作为状态让智能体感知到运动信息。其次奖励需要截断到 ([-1, 1]) 区间避免不同游戏奖励量纲差异对训练造成过大冲击。再者优化器建议从Adam换成RMSProp学习率调整到2.5e-4这是DeepMind论文中验证过的稳定配置。最后训练步数要扩大到百万级别我用CartPole几十步就能看到效果但在Atari上几万步基本只是热身。我的建议是不要直接上来就挑战Atari。先在CartPole上把DQN的代码框架、调参手感和调试技巧练熟再逐步往图像输入、连续状态空间这些方向延展每一步都确认当前环节稳定了再往前走这样效率其实更高。5. 写在最后的几点实践心得DQN这套代码我前前后后重写过很多遍每次重写都会对强化学习本身有一层新的理解。它最值得学习的地方不在于多么复杂的网络结构而在于那套应对非平稳目标、样本相关性的精巧设计。这些思路在后续的PPO、SAC等算法中也能看到延续的影子基础打牢了后面学什么都快。最后分享两个调参时的土办法。第一个是画图大法把loss和reward曲线都画出来看得见的变化比任何日志都直观。第二个是单步调试法拿一个最小的环境把每一步的前向计算、目标构造、反向更新都打印出来亲手算一遍数值比读十遍论文都管用。这些办法听上去笨拙但对付试验性代码异常管用。
返回列表