
1. 从CS285的作业说起Q-learning到底在解决什么问题如果你正在啃CS285深度强化学习的课程大概率会在某个深夜对着Q-learning的代码发呆——明明理论推导看懂了公式也背下来了但一上手写代码就发现怎么跟想象的不一样收敛慢、震荡、过估计、样本效率低这些问题一个接一个冒出来。我当初第一次实现Q-learning的时候在CartPole上跑了整整一个下午都没收敛后来才发现是学习率设大了导致Q值直接发散。所以这篇内容我想从一个实际写过、调过、踩过坑的人的角度把Q-learning这件事从头到尾讲清楚。Q-learning本质上是一种基于值的强化学习算法它的核心目标是学到一个函数Q(s, a)告诉你“在状态s下做动作a未来能拿到多少累计回报”。一旦这个函数学好了策略就很简单在每个状态选Q值最大的那个动作就行。听起来很直白对吧但魔鬼藏在细节里。CS285的课程之所以把Q-learning放在比较靠前的位置是因为它承上启下——往前连接了动态规划的思路往后引出了DQN、Double Q-learning、Dueling Network等一系列深度强化学习的核心方法。这篇文章适合谁看如果你正在学CS285或者任何一门强化学习课程正在被Q-learning的作业折磨如果你已经看过理论但不知道怎么落地写代码如果你写出来了但效果不好不知道怎么调——那这篇内容就是为你准备的。我会从算法设计思路、核心公式推导、实操步骤、参数选择、常见问题排查这几个维度把Q-learning彻底拆开讲透。不会只给你一堆公式而是告诉你每个公式背后的直觉是什么代码里每一行在干什么以及我实际踩过的那些坑。2. Q-learning的整体设计思路与方案选型2.1 为什么是“基于值”而不是“基于策略”强化学习的算法大致分两派一派是基于策略的Policy-based直接学一个策略函数π(a|s)告诉你每个状态下该做什么动作另一派是基于值的Value-based学一个值函数间接推导出策略。Q-learning属于后者。为什么CS285要先讲基于值的方法我的理解是基于值的方法在离散动作空间下更直观、更容易实现而且它的核心思想——用贝尔曼方程做迭代更新——是整个强化学习的基石。你把这个搞懂了后面学Policy Gradient、Actor-Critic都会轻松很多。具体来说Q-learning维护一个Q表或者Q网络记录每个状态-动作对的估值。每次执行一个动作、观察到奖励和下一个状态之后用**时序差分TD的方式更新Q值。这个过程不需要知道环境的转移概率也就是不需要model所以它是无模型的model-free**方法。这一点很关键——现实中大部分场景你根本不知道环境的动力学模型所以model-free的方法适用范围更广。2.2 贝尔曼最优方程Q-learning的理论根基Q-learning的更新公式来源于贝尔曼最优方程。我先把这个方程写出来然后用人话解释它。Q*(s, a) E[r γ · max_{a} Q*(s, a)]翻译一下在状态s做动作a的最优价值等于“即时奖励r”加上“折扣因子γ乘以下一个状态s下所有动作中最大的Q值”。这个“max”就是Q-learning的灵魂——它假设你在下一步会做出最优选择所以用最大值来更新当前估值。为什么这个方程成立因为它定义的是最优价值函数。你不断用这个方程去迭代更新Q值最终Q会收敛到Q*。这就像你在一个迷宫里不断试错每次走到一个新位置就回头修正之前对各个岔路口的评价走得多了每个路口该往哪走就越来越清晰。2.3 On-policy还是Off-policyQ-learning的独特定位这里有一个很多人初学时会混淆的点Q-learning是off-policy的算法。什么意思就是说它用来更新Q值的数据可以不是当前策略产生的。你完全可以用一个随机策略去探索环境收集到的数据照样能用来更新Q值而且最终学到的还是最优策略。对比一下SARSA——SARSA是on-policy的它的更新用的是实际执行的下一个动作a的Q值而不是max。这个区别看起来很小但影响很大。Q-learning更激进因为它总是假设下一步选最优动作SARSA更保守因为它考虑的是实际策略的行为。在悬崖行走Cliff Walking这个经典环境里Q-learning会学到贴着悬崖走的最优路径而SARSA会学到绕远但更安全的路径。理解这个区别对你后面调参和选择算法非常重要。2.4 从表格到神经网络为什么需要函数逼近最原始的Q-learning用一个表格存Q值状态数少的时候没问题。但CS285的作业里状态空间往往是连续的比如CartPole的观测是4维连续值你不可能枚举所有状态。这时候就需要函数逼近——用一个神经网络来拟合Q(s, a)这个映射。这就引出了DQNDeep Q-Network的核心思路。但直接用神经网络做Q-learning会遇到几个大问题样本之间高度相关、目标值不断变化导致训练不稳定。解决方案是经验回放Experience Replay和目标网络Target Network。这两个技巧在CS285的作业里都会涉及后面我会详细讲怎么实现。3. 核心细节解析与实操要点3.1 Q-learning更新公式的逐步拆解先把更新公式完整写出来Q(s, a) ← Q(s, a) α · [r γ · max_{a} Q(s, a) - Q(s, a)]这里面有几个关键组成部分我逐个拆解TD目标Targetr γ · max_{a} Q(s, a)。这是你“期望”Q值应该达到的水平。它由即时奖励和下一步的最优估值组成。TD误差TD ErrorTD目标减去当前Q值。这个差值告诉你当前的估计偏了多少。如果TD误差是正的说明实际回报比预期好应该往上调Q值反之往下调。学习率α控制每次更新幅度。α太大Q值会震荡甚至发散α太小收敛速度慢得让人抓狂。我一般从0.001开始试如果训练曲线震荡就降到0.0005如果太慢就升到0.005。折扣因子γ决定你有多看重未来奖励。γ0就是完全短视只看眼前γ接近1就是非常有远见。大部分任务用0.95到0.99之间比较合适。CartPole这种任务γ0.99效果不错。3.2 探索与利用的平衡ε-greedy策略Q-learning有一个天然的问题如果你总是选当前Q值最大的动作那你就永远不会去尝试那些可能更好的动作。这就是**探索与利用Exploration vs Exploitation**的经典困境。最常用的解决方案是ε-greedy策略以ε的概率随机选一个动作探索以1-ε的概率选Q值最大的动作利用。ε通常从一个较大的值比如1.0开始随着训练逐渐衰减到一个较小的值比如0.05或0.01。我实际用下来线性衰减比指数衰减更可控。比如你计划训练10000步可以让ε从1.0线性降到0.05。这样前期充分探索后期稳定利用。但要注意ε不能降到0否则一旦环境有变化智能体就完全失去适应能力了。注意有些实现里会用ε的最小值作为下限比如ε_min0.05即使衰减到底也保留5%的随机性。这个细节在CS285的作业里经常被忽略但对最终效果影响不小。3.3 经验回放缓冲区的设计与实现经验回放是DQN能训练成功的关键组件之一。它的核心思想很简单把智能体与环境交互产生的经验(s, a, r, s, done)存到一个缓冲区里训练的时候从中随机采样一批数据。为什么要随机采样因为连续的经验之间高度相关——你在CartPole里连续几步的状态几乎一样。如果直接用这些数据训练神经网络会“记住”最近的经验而忘记之前的导致灾难性遗忘。随机采样打破了这种相关性让每次更新都基于一个更多样化的数据分布。缓冲区的容量一般设10万到100万条经验。太小了容易过拟合最近的数据太大了会占用大量内存而且可能包含太多过时的策略产生的数据。我一般用10万条起步如果任务复杂就加到50万。采样的时候用均匀随机采样就行不需要按优先级采样那是Prioritized Experience Replay做的事CS285后面会讲。每次采样一个batchbatch size通常设32到256之间。太小了梯度噪声大太大了计算慢而且可能降低样本多样性。3.4 目标网络为什么需要以及怎么用目标网络是另一个让Q-learning能work的关键技巧。问题出在TD目标上r γ · max_{a} Q(s, a)。如果你用同一个网络既算当前Q值又算TD目标那目标本身在每次更新后都会变就像你在追一个自己也在跑的目标训练极不稳定。解决方案是维护两个网络**在线网络Online Network**用来选动作和计算当前Q值**目标网络Target Network**用来计算TD目标。目标网络的参数不是每次更新都变而是每隔C步从在线网络复制一次硬更新或者每次只缓慢更新一点点软更新。硬更新的C一般设1000到10000步。软更新的更新率τ一般设0.001到0.01。我个人的经验是软更新在小任务上更稳硬更新在大任务上更常见。CS285的作业里两种都可能要求实现建议都试一下。4. 实操过程与核心环节实现4.1 环境搭建与依赖安装先说一下我用的环境配置。Python 3.8以上PyTorch 1.10以上gym 0.21以上。如果你用的是CS285的官方代码框架它已经帮你封装好了环境接口你只需要实现算法部分。但如果你想从零写一遍下面这些依赖是必须的pip install gym torch numpy matplotlibgym提供环境torch用来搭网络numpy做数值计算matplotlib画训练曲线。如果你要用MuJoCo的环境CS285后半段会用到还需要额外安装mujoco-py那个配置比较麻烦建议先用CartPole把Q-learning跑通再说。4.2 Q表的实现离散状态版本先从最简单的Q表版本开始这能帮你理解算法核心。假设状态空间是离散的比如FrozenLake这种网格世界。import numpy as np class QLearningAgent: def __init__(self, n_states, n_actions, alpha0.001, gamma0.99, epsilon1.0): self.Q np.zeros((n_states, n_actions)) self.alpha alpha self.gamma gamma self.epsilon epsilon self.n_actions n_actions def select_action(self, state): if np.random.random() self.epsilon: return np.random.randint(self.n_actions) return np.argmax(self.Q[state]) def update(self, state, action, reward, next_state, done): td_target reward if not done: td_target self.gamma * np.max(self.Q[next_state]) td_error td_target - self.Q[state, action] self.Q[state, action] self.alpha * td_error这段代码就是Q-learning最核心的逻辑。注意done的处理——如果回合结束了就没有下一个状态TD目标就是即时奖励本身。这个细节很多人会漏掉导致Q值在终止状态附近被高估。4.3 神经网络版本的实现连续状态版本连续状态就需要用神经网络了。下面是我在CartPole上用的网络结构import torch import torch.nn as nn class QNetwork(nn.Module): def __init__(self, obs_dim, n_actions, hidden_dim64): super().__init__() self.net nn.Sequential( nn.Linear(obs_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, n_actions) ) def forward(self, x): return self.net(x)CartPole的观测是4维动作是2个左推或右推所以obs_dim4n_actions2。隐藏层用64个神经元两层就够了。太深的网络在小任务上反而容易过拟合而且训练慢。训练循环的核心逻辑for episode in range(n_episodes): state env.reset() done False while not done: action agent.select_action(state) next_state, reward, done, _ env.step(action) replay_buffer.push(state, action, reward, next_state, done) if len(replay_buffer) batch_size: batch replay_buffer.sample(batch_size) agent.update(batch) state next_state # 每隔C步更新目标网络 if episode % target_update_freq 0: target_net.load_state_dict(online_net.state_dict())4.4 参数选择与调优记录我在CartPole上实际调参的记录供你参考参数初始尝试最终采用说明学习率α0.010.0010.01太大导致Q值发散折扣因子γ0.90.990.9学到的策略太短视ε初始值0.51.0从1.0开始探索更充分ε衰减步数500010000衰减太快导致探索不足缓冲区大小10000100000太小导致样本相关性高Batch size6412864梯度噪声太大目标网络更新频率100500100太频繁目标不稳定隐藏层大小326432表达能力不够这张表是我踩了无数坑之后总结出来的。你刚开始的时候可以直接用“最终采用”这一列的参数应该能在CartPole上跑到200步左右的平均回报满分是200。4.5 训练过程监控与可视化训练过程中一定要监控几个关键指标平均回报、平均Q值、TD误差、ε值。平均回报告诉你策略好不好平均Q值告诉你估值有没有发散TD误差告诉你学习是否稳定ε值告诉你探索程度。我一般每100个回合画一次平均回报曲线。如果曲线震荡得很厉害说明学习率太大或者batch size太小。如果曲线一直不涨可能是探索不够ε太大或者网络容量不够。如果Q值突然变得很大那就是发散了赶紧降低学习率。import matplotlib.pyplot as plt def plot_training(returns, q_values): fig, axes plt.subplots(1, 2, figsize(12, 4)) axes[0].plot(returns) axes[0].set_title(Average Return) axes[1].plot(q_values) axes[1].set_title(Average Q Value) plt.savefig(training_curve.png)5. 常见问题与排查技巧实录5.1 Q值发散或爆炸怎么办这是Q-learning最常见的问题没有之一。症状是训练过程中Q值越来越大最后变成NaN。原因通常有三个学习率太大、网络太大、或者没有用目标网络。排查顺序先把学习率降一个数量级试试比如从0.001降到0.0001。如果还不行检查目标网络有没有正确更新——很多人忘了复制参数或者复制频率设错了。再不行就减小网络规模把隐藏层从128降到64甚至32。还有一个容易被忽略的原因奖励尺度。如果你的奖励值很大比如100、1000TD目标也会很大容易导致梯度爆炸。解决方案是把奖励归一化到[-1, 1]或者[0, 1]之间。这个技巧在CS285的作业里经常被用到。5.2 训练不收敛或收敛到次优策略如果Q值没有发散但策略就是学不好可能的原因包括探索不足、经验回放有问题、或者网络结构不合适。探索不足的典型表现是智能体反复做同一个动作。检查ε的衰减曲线确保在训练早期有足够的随机性。我一般会确保前10%的训练步数里ε保持在0.5以上。经验回放的问题通常是缓冲区太小或者采样策略有问题。如果你用的是均匀采样确保缓冲区里至少有几千条经验再开始训练。如果缓冲区里全是最近的数据那跟不用回放差不多。网络结构的问题比较隐蔽。如果任务的状态空间维度很高比如图像输入两层MLP肯定不够需要上卷积网络。如果动作空间很大比如连续动作Q-learning就不太适用了得换DDPG或者SAC。5.3 过估计问题与Double Q-learningQ-learning有一个固有的缺陷过估计Overestimation。因为TD目标里用的是max操作而Q值本身是有噪声的max会倾向于选到那些被高估的动作导致Q值系统性偏高。解决方案是Double Q-learning用两个网络一个负责选动作一个负责估值。具体来说用在线网络选出Q值最大的动作a*然后用目标网络计算Q(s, a*)作为TD目标。这样就把“选动作”和“估价值”分开了有效减少过估计。在CS285的作业里Double Q-learning通常作为一个改进项出现。实现起来很简单就是把TD目标的计算从target_net(next_state).max()改成target_net(next_state)[online_net(next_state).argmax()]。改动很小但效果提升明显。5.4 常见问题速查表症状可能原因解决方案Q值发散/NaN学习率太大降低学习率10倍Q值发散/NaN奖励尺度太大归一化奖励策略不收敛探索不足增大ε或减慢衰减策略不收敛缓冲区太小增大缓冲区到10万策略震荡batch size太小增大到128或256收敛到次优过估计使用Double Q-learning训练太慢网络太大减小隐藏层训练太慢目标网络更新太频繁增大更新间隔回报突然下降灾难性遗忘检查经验回放是否正常工作回报一直不涨学习率太小适当增大学习率5.5 实操心得与避坑建议最后分享几个我在实际写Q-learning时总结的经验都是文档里不会写的第一先在小任务上验证正确性。不要一上来就跑复杂的任务。先用FrozenLake或者CartPole这种简单环境把算法跑通确认Q值收敛、策略正确再去挑战更难的环境。我见过太多人直接上Atari结果调了一周都不知道问题出在哪。第二随机种子很重要。强化学习的方差很大同一个算法不同的随机种子可能结果差很多。建议至少跑3到5个种子取平均结果。如果只有一个种子的结果好那很可能是运气。第三训练曲线比最终结果更重要。不要只看最后跑了多少分要看整个训练过程的曲线。如果曲线是稳步上升的说明算法在正常学习如果曲线剧烈震荡或者突然下降说明有问题需要排查。第四善用TensorBoard。把回报、Q值、TD误差、ε值都记录下来可视化之后很多问题一眼就能看出来。比如Q值曲线如果一直往上飘那就是过估计或者发散的前兆。第五不要过度调参。强化学习的参数确实多但大部分情况下默认参数就能work。如果调了好几次都不行可能是算法实现有bug而不是参数问题。这时候应该回去检查代码逻辑特别是TD目标的计算和梯度更新的部分。Q-learning这个算法理论看起来简单但真正写好、调好需要不少经验。CS285的作业设计得很好它会逼着你把每个细节都搞清楚。我当初做的时候光是目标网络的更新逻辑就改了三遍才跑通。但一旦你把这个算法吃透了后面学DQN的各种变体、学Actor-Critic、学Policy Gradient都会觉得顺理成章。希望这篇内容能帮你少走一些弯路把Q-learning真正搞明白。