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

资讯详情

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

CS285 Q-learning实战:从Bellman最优方程到代码实现与调参

CS285 Q-learning实战:从Bellman最优方程到代码实现与调参 1. 从CS285的作业说起Q-learning到底在解决什么问题如果你正在跟CS285深度强化学习这门课大概率会在前几讲就被Q-learning绕得有点晕。我第一次接触它的时候脑子里全是问号这玩意儿和策略梯度到底啥关系为什么一会儿要学Q函数一会儿又要取max那个Bellman方程怎么推着推着就变成了一个迭代更新的代码后来把作业从头到尾写了一遍又拿几个小环境反复调参才算真正把这条线捋清楚。Q-learning的核心目标其实非常朴素在不知道环境模型也就是不知道状态转移概率和奖励函数的前提下学会一个动作价值函数Q(s, a)然后每次选动作时挑Q值最大的那个。它属于value-based方法里最经典的一支和policy gradient那种直接优化策略参数的思路是两条路。CS285把它放在前面讲是因为它足够基础又足够能暴露强化学习里几个核心难点时序差分、自举bootstrapping、off-policy、探索与利用的平衡。这篇文章我打算按我自己做作业和跑实验的顺序来写不搞教科书式的定义堆砌。我会把Q-learning在CS285语境下的完整实现链路拆开从Bellman最优算子的推导直觉到target network为什么要加到replay buffer怎么设计再到实际训练时那些让人抓狂的震荡和发散问题。适合正在做CS285作业的同学也适合任何想从代码层面真正搞懂Q-learning的人。你不需要先精通策略梯度但最好对马尔可夫决策过程MDP的基本符号有点概念不然看到后面会有点吃力。2. Q-learning的核心原理拆解为什么是“最优”而不是“当前”2.1 从Bellman期望方程到Bellman最优方程先把这个最容易被跳过、但后面所有代码都依赖的推导说清楚。对于一个固定策略π动作价值函数满足Bellman期望方程Q^π(s, a) r(s, a) γ · E_{s~P}[ V^π(s) ]其中V^π(s) E_{a~π}[ Q^π(s, a) ]。注意这里对下一个动作a是按策略π的分布求期望因为你在评估的是“如果我一直按π走能拿多少回报”。Q-learning干的事情不一样。它不评估某个固定策略而是直接逼近最优动作价值函数Q*。对应的方程是Bellman最优方程Q*(s, a) r(s, a) γ · E_{s~P}[ max_{a} Q*(s, a) ]差别就在那个max。它表达的意思是在下一个状态我会选当前认为最好的动作而不是按某个策略采样。这个max是Q-learning的灵魂也是它off-policy能力的来源——因为不管我实际用的是什么行为策略比如ε-greedy我更新时用的都是max所以我学的是最优Q而不是行为策略的Q。提示很多同学第一次写代码时会把target写成Q(s, a)其中a是从当前策略采样的那就变成了SARSA不是Q-learning。这个区别在作业里经常是扣分点。2.2 时序差分更新用估计去更新估计有了最优方程理论上可以解一个非线性方程组得到Q*。但状态空间一大就没法解。于是我们用随机逼近每次拿到一个transition (s, a, r, s)就朝目标方向挪一小步Q(s, a) ← Q(s, a) α · [ r γ · max_{a} Q(s, a) − Q(s, a) ]方括号里那坨叫TD误差。它衡量的是“我原来的估计”和“用一步真实奖励加下一步估计拼出来的新目标”之间的差距。这里有个关键点目标里用了Q(s, a)而Q本身还在更新这就是自举bootstrapping。好处是不需要等整个episode结束就能学方差比蒙特卡洛小坏处是估计偏差会传播可能不稳定。α是学习率。在CS285的作业里通常用Adam优化器代替手工α但理解上还是这个形式。γ是折扣因子控制你多看重未来。γ越接近1越看重长期回报但自举的误差也会被放大更多步训练更容易飘。2.3 为什么Q-learning是off-policy的这一点值得单独拎出来。off-policy的意思是行为策略用来采样数据的策略和目标策略正在学习评估的策略可以不同。Q-learning的更新目标是max Q它不关心这个(s, a, r, s)是哪个策略产生的。哪怕数据是随机策略采的甚至是几年前旧策略采的只要放进replay buffer照样能用来更新Q*。这个性质在实际中非常值钱。你可以用高探索性的策略去采集大量数据然后反复从buffer里采样训练样本效率比on-policy的policy gradient高很多。但代价是如果行为策略覆盖的状态动作空间不够Q函数在某些区域就没有数据支撑max操作会把这些区域的Q值高估进而影响策略。这就是后来Double Q-learning、保守Q学习等一系列工作的动机。3. 在CS285里落地Q-learning网络、Buffer与训练循环3.1 函数逼近器的选择与网络结构CS285的Q-learning作业通常要求用神经网络来参数化Q(s, a)。有两种常见结构状态输入、输出所有动作的Q值网络输入是状态s输出维度等于动作数|A|每个分量是Q(s, a_i)。这种结构适合离散动作空间前向一次就能拿到所有动作的Q取max非常方便。状态动作拼接输入、输出单个标量输入是(s, a)拼接输出一个Q值。这种适合连续动作但取max就需要优化离散情况下效率低。作业里大多是离散动作比如CartPole所以第一种更常见。网络本身不用太深两层全连接、每层64到256个单元、ReLU激活基本就能跑。我试过用更宽的网络在简单环境里反而更容易过拟合到早期数据训练曲线更抖。注意输出层不要加激活函数。Q值可以是任意实数加了ReLU或者tanh会把值域限制住导致无法拟合真实回报。这个坑我踩过当时loss一直下不去查了半天才发现是输出层多写了个激活。3.2 Replay Buffer的设计细节Replay buffer是off-policy方法的关键组件。它就是一个固定容量的队列存(s, a, r, s, done)五元组。新数据进来旧数据出去。训练时从里面随机采样一个batch。几个实操要点容量太小则数据相关性太强训练不稳太大则旧数据太多和当前策略分布偏离太远。CS285作业里通常给10万到100万。CartPole这种简单任务5万到10万就够。采样必须均匀随机采样不能按时间顺序取。按顺序取等于变相on-policy会破坏off-policy的样本效率优势还容易过拟合最近的数据。done的处理如果s是终止状态target就只是r不加γ·max Q(s, a)。这个细节写错的话终止状态的Q会被系统性高估。我自己的习惯是在buffer里存done标志计算target时用(1 - done)乘以后续项。这样一行代码就能处理不容易漏。3.3 Target Network为什么需要一份“慢半拍”的副本如果直接用同一个网络算target和预测每次更新都会同时改变两边目标在追着自己跑很容易发散。Target network的思路是复制一份Q网络参数冻结每隔C步才把主网络的参数同步过来。这样在一段时间内target是稳定的训练更像是在拟合一个固定目标。同步方式有两种硬更新每隔C步直接复制参数。简单但同步瞬间会有跳变。软更新Polyak每次更新时θ_target ← τ·θ (1−τ)·θ_targetτ通常取0.001到0.01。更平滑DDPG、SAC这些都用软更新。CS285的Q-learning作业里两种都可能出现看具体版本。我实测下来简单离散任务硬更新C1000左右就够软更新τ0.005也很稳。如果训练曲线周期性抖动多半是硬更新的同步周期设得太短。3.4 完整训练循环的伪代码与逐行解释把上面几块拼起来一个标准的Q-learning训练循环长这样for step in range(total_steps): # 1. 用当前Q网络加ε-greedy选动作 if random() epsilon: a env.action_space.sample() else: a argmax(Q_net(s)) # 2. 执行动作拿到transition s_next, r, done, _ env.step(a) buffer.add(s, a, r, s_next, done) s s_next # 3. 从buffer采样一个batch batch buffer.sample(batch_size) s_b, a_b, r_b, s_next_b, done_b batch # 4. 计算target with torch.no_grad(): q_next target_net(s_next_b).max(dim1)[0] target r_b gamma * (1 - done_b) * q_next # 5. 计算当前Q并做梯度下降 q_pred Q_net(s_b).gather(1, a_b) loss mse_loss(q_pred, target) optimizer.zero_grad() loss.backward() optimizer.step() # 6. 定期同步target network if step % target_update_freq 0: target_net.load_state_dict(Q_net.state_dict())逐行看几个容易出错的点第1步的ε-greedyε通常从1.0线性退火到0.05或0.1。前期多探索后期多利用。如果ε一直很大Q学不准一直很小可能陷入局部。第4步的torch.no_grad()必须加否则target会带着梯度既浪费显存又可能污染计算图。第5步的gather是按采样的动作索引取出对应的Q值。如果动作是one-hot也可以用逐元素相乘再求和。第6步的同步频率很关键前面说过。4. 实操中那些让人抓狂的问题与排查思路4.1 Q值爆炸与发散最常见的现象是loss突然变成NaN或者Q值涨到几百上千。原因通常有几个学习率太大自举本身就会放大误差学习率再大直接飞。先降到1e-4甚至1e-5试试。target network没生效检查同步逻辑是不是每步都在同步等于没有target。奖励尺度太大如果环境奖励是几百的量级Q值自然大。可以考虑对奖励做缩放或者调小学习率。没有梯度裁剪加个clip_grad_norm_(params, 10)能救很多情况。我遇到过一次Q值在几千步后突然爆炸查下来是target network的同步频率写成了每步同步等于完全没有稳定目标。改成每1000步同步后曲线立刻正常。4.2 训练曲线震荡不收敛震荡通常意味着target变化太快或者数据分布变化太剧烈。可以尝试增大target更新周期或者改用软更新。减小学习率。增大batch size降低梯度估计方差。检查replay buffer是不是太小导致采样数据高度相关。4.3 学到的策略很“怂”或者很“莽”这往往和ε、γ、奖励设计有关。γ太小智能体只看眼前表现得很短视γ太大又容易被远期噪声干扰。ε退火太快探索不足策略可能卡在次优。我的经验是先把ε退火拉长一点观察Q值是否在合理范围内增长再逐步收紧。4.4 常见问题速查表现象可能原因排查方向loss变NaN学习率过大、无梯度裁剪降lr、加clipQ值持续增大不收敛target未冻结、奖励尺度大检查同步逻辑、缩放奖励策略不改进ε太小、探索不足增大ε或延长退火训练初期就崩网络输出层有激活去掉输出层激活曲线周期性抖动硬更新周期太短增大C或改软更新终止状态Q值异常done未处理用(1-done)乘后续项提示排查时优先怀疑target network和done处理这两个地方出错最隐蔽也最致命。5. 从Q-learning延伸出去它和后续方法的关系把Q-learning写通之后再看CS285后面的内容会顺很多。Double Q-learning解决的是max导致的高估问题做法是用两个网络一个选动作一个评估。Dueling Q-learning把Q拆成状态价值和优势函数让网络在不同状态下更高效地学习。优先经验回放Prioritized Replay则是对TD误差大的样本加大采样概率提升样本效率。再往后到连续动作空间直接取max就不可行了于是有了DDPG用actor网络近似argmax、SAC用熵正则鼓励探索。这些方法的核心骨架还是Q-learning那一套TD target、target network、replay buffer。所以把这一讲吃透后面基本是换零件而不是换发动机。我自己在做完Q-learning作业后回头把target network和replay buffer单独抽出来做了几个消融实验发现去掉target network后CartPole几乎学不起来去掉replay buffer后样本效率掉一大截。这种亲手验证比看十遍公式都管用。如果你时间允许强烈建议也做一遍消融你会对每个组件的必要性有完全不同的体感。
返回列表