📑 本页目录(点开跳转)
08 · DQN:状态存不下了怎么办
⏱ 30 分钟 | ⭐ 深度强化学习的起点
🎯 一句话
把 Q 表换成神经网络。 听起来只是个小改动,但它会让前面所有的收敛保证全部失效 —— 这一章讲那三个救命的补丁。
💥 一、为什么必须换掉表格
Q 表的大小 = 状态数 × 动作数
格子迷宫 10×10: 100 × 4 = 400 个格子 ✅ 存得下
Atari 一帧画面: 256^(84×84×4) 个状态 💀 宇宙里没这么多原子
而且表格还有一个更根本的问题:
⭐ 表格法没有【泛化】能力
学会了「敌人在左边 30 像素时该往右躲」
→ 遇到「敌人在左边 31 像素」,表格认为这是【完全不同】的状态
→ 之前学的一点用都没有 💀
神经网络天然会泛化:相似的输入 → 相似的输出 ✅
🔑 这才是用神经网络的真正理由。 不只是"存不下",是"每个状态都要单独学一遍"根本不现实。
🧠 二、DQN 的想法:用网络逼近 Q
$$Q(s,a;\theta) \approx Q^*(s,a)$$
💡 一个小设计但很重要:输出所有动作的 Q,而不是把 (s,a) 一起输入。 这样选动作只需要一次前向传播,而不是每个动作跑一遍。
损失函数(直接来自第 6 章的 Q-learning):
$$L(\theta) = \Big(\underbrace{r + \gamma\max_{a'}Q(s',a';\theta^-)}_{\text{目标}} - Q(s,a;\theta)\Big)^2$$
⭐ 看这个式子:它就是把 Q-learning 的更新式写成了均方误差。 于是强化学习问题变成了一个监督学习问题 —— 可以用梯度下降了。
💀 三、直接这么做会崩:三个致命问题
问题 1:样本高度相关
监督学习假设:样本独立同分布(i.i.d.)
强化学习现实:连续的帧几乎一模一样 💀
→ 梯度方向高度相关 → 训练极不稳定
→ 相当于拿同一个样本连着更新一百次
问题 2:目标在动(追自己的尾巴)⭐
L = (r + γ max Q(s',a';θ) − Q(s,a;θ))²
↑ ↑
目标用 θ 预测也用 θ
→ 你更新 θ 去逼近目标
→ 目标因为 θ 变了,也跟着动了 💀
→ 像追自己的尾巴,可能永远追不上,甚至发散
问题 3:数据分布随策略改变
策略变好 → 去的状态变了 → 训练数据分布变了
→ 之前学的可能失效(灾难性遗忘)
🔑 这三个合起来叫「致命三要素」(deadly triad): 函数近似 + 自举(bootstrapping)+ 离策略 —— 三者同时出现时,理论上可以发散。 ⭐ 第 6 章那个"表格 Q-learning 保证收敛"的定理,在这里彻底失效了。
🩹 四、三个救命补丁
补丁 1:经验回放(Experience Replay)⭐
把经历过的 (s, a, r, s', done) 存进一个大池子(比如 100 万条)
训练时【随机抽一批】出来,而不是用刚发生的
✅ 打破了样本相关性
✅ 一条经验可以被【重复利用】很多次 → 样本效率大幅提升
✅ 平滑了数据分布的变化
⭐ 注意这一步只有 off-policy 算法能做 —— 池子里是旧策略采的数据。 🔗 这就是第 6 章说"off-policy 是后面一切基础"的第一次兑现。
补丁 2:目标网络(Target Network)⭐
💡 C 怎么设:常用 1000~10000 步。 太小 → 还是在追尾巴;太大 → 目标太陈旧,学得慢。
补丁 3:奖励裁剪与预处理
· 奖励裁剪到 [-1, 1] → 不同游戏的分数量级差异巨大,不裁剪没法用同一套超参
· 画面转灰度、缩到 84×84、堆叠 4 帧 → 降维 + 恢复马尔可夫性(第 2 章)
🔨 五、核心代码
import torch, torch.nn as nn, numpy as np, random
from collections import deque
class DQN:
def __init__(self, net, n_actions, gamma=0.99, lr=1e-4, buf=100_000):
self.q, self.q_target = net, __import__('copy').deepcopy(net)
self.opt = torch.optim.Adam(self.q.parameters(), lr=lr)
self.buffer = deque(maxlen=buf) # ⭐ 补丁1:经验回放池
self.gamma, self.n_actions = gamma, n_actions
self.step_count = 0
def act(self, s, eps):
if random.random() < eps:
return random.randrange(self.n_actions)
with torch.no_grad():
return int(self.q(torch.as_tensor(s).unsqueeze(0)).argmax())
def learn(self, batch_size=32, sync_every=1000):
if len(self.buffer) < batch_size:
return
s, a, r, s2, d = map(np.array, zip(*random.sample(self.buffer, batch_size)))
s, s2 = torch.as_tensor(s).float(), torch.as_tensor(s2).float()
a, r, d = torch.as_tensor(a), torch.as_tensor(r).float(), torch.as_tensor(d).float()
q = self.q(s).gather(1, a.unsqueeze(1)).squeeze(1)
with torch.no_grad(): # ⭐ 补丁2:目标网络,且不回传梯度
tgt = r + self.gamma * self.q_target(s2).max(1)[0] * (1 - d)
loss = nn.functional.smooth_l1_loss(q, tgt) # ⭐ Huber,比 MSE 抗离群
self.opt.zero_grad(); loss.backward()
nn.utils.clip_grad_norm_(self.q.parameters(), 10) # ⭐ 梯度裁剪
self.opt.step()
self.step_count += 1
if self.step_count % sync_every == 0:
self.q_target.load_state_dict(self.q.state_dict())
⚠️ 四个最容易写错的地方:
| 坑 | 后果 |
|---|---|
目标里忘了 torch.no_grad() |
梯度回传到目标网络,训练直接乱掉 ⭐ |
忘了 * (1 - d) |
终止态算上了不存在的未来价值(和第 5 章同一个坑) |
| 用 MSE 而不是 Huber | TD 误差偶尔很大时梯度爆炸 |
| 目标网络忘了同步 / 同步太频繁 | 前者目标永远陈旧,后者等于没加这个补丁 |
🚀 六、后续改进(知道名字和一句话就够)
| 改进 | 解决什么 | 一句话 |
|---|---|---|
| Double DQN ⭐ | 最大化偏差(第 6 章) | 用主网络选动作、目标网络评估 |
| Dueling DQN | 有些状态下动作无所谓 | 拆成 V(s) + 优势 A(s,a) |
| 优先经验回放 PER ⭐ | 均匀采样浪费 | TD 误差大的样本多抽 —— 错得多的地方多学 |
| Noisy Nets | ε-贪心探索太笨 | 给权重加噪声(第 7 章) |
| Rainbow | —— | 把上面全部组合起来,效果最好 |
💡 Double DQN 的一行改动(值得记住):
# 原始 DQN: 用目标网络【既选又评】→ 高估
tgt = r + gamma * self.q_target(s2).max(1)[0] * (1 - d)
# Double DQN:主网络【选】,目标网络【评】⭐
a2 = self.q(s2).argmax(1, keepdim=True)
tgt = r + gamma * self.q_target(s2).gather(1, a2).squeeze(1) * (1 - d)
🧱 七、DQN 的天花板
| 限制 | 说明 |
|---|---|
| 只能处理离散动作 ⭐ | max_a 要遍历所有动作。动作连续就完蛋 |
| 样本效率仍然低 | Atari 要几千万帧 |
| 对超参数敏感 | 学习率、C、buffer 大小都很关键 |
⭐ 第一条就是第 9 章存在的理由: 「方向盘转多少度」「机器人关节用多大力」「生成哪个 token(词表 5 万)」—— 这些都没法对所有动作取 max。 → 换一条路:不估价值了,直接优化策略。
🔗 八、和站内其他章的关系
| 相关的地方 | 这里的位置 |
|---|---|
| 第 6 章 Q-learning | DQN = 它 + 神经网络 + 三个补丁 |
| 第 6 章 off-policy | 经验回放的前提 ⭐ |
| 第 2 章堆叠 4 帧 | 补丁 3 里的预处理 |
| ML基础 12 CNN | DQN 的网络主体 |
| ML基础 13 梯度裁剪 | ⭐ 同一个技巧:那里讲了梯度为什么会爆(连乘次数长、乘的还是同一个矩阵)和 max_norm 该取多少——DQN 里裁的是同一个东西 |
✅ 检查点
- 换掉表格的两个理由是什么?哪个更根本?
- DQN 的损失函数是什么?它把 RL 变成了什么问题?
- 「致命三要素」是哪三个?它意味着什么?
- 经验回放解决什么问题?为什么只有 off-policy 能用?
- 目标网络解决什么问题?C 太小太大分别会怎样?
- 代码里忘了
torch.no_grad()会怎样?忘了*(1-d)呢? - Double DQN 的一行改动是什么?解决什么?
- DQN 的最大天花板是什么?为什么它逼出了第 9 章?
👀 答案
- ①状态太多存不下(Atari 是 256^(84×84×4))②表格没有泛化能力——"敌人在左30像素"和"左31像素"被当成完全不同的状态。第二个更根本:不只是存不下,是每个状态都要单独学一遍不现实。
L = (r + γ max Q(s',a';θ⁻) − Q(s,a;θ))²。它把 RL 变成了监督学习问题(拟合一个回归目标),于是能用梯度下降。- 函数近似 + 自举 + 离策略。三者同时出现时理论上可以发散——第 6 章的表格收敛保证在这里彻底失效。
- 解决样本高度相关(连续帧几乎一样,梯度方向高度相关)。只有 off-policy 能用因为池子里是旧策略采的数据。
- 解决目标在动、追自己尾巴的问题。C 太小还是在追尾巴;C 太大目标太陈旧学得慢(常用 1000~10000)。
- 忘
no_grad()→ 梯度回传到目标网络,训练直接乱掉;忘*(1-d)→ 终止态算进了不存在的未来价值。 - 主网络选动作、目标网络评估:
a2=q(s2).argmax(); tgt=r+γ·q_target(s2).gather(1,a2)。解决最大化偏差(max 系统性挑中被高估的动作)。 - 只能处理离散动作——
max_a要遍历所有动作,动作连续就没法算。所以第 9 章换路:不估价值,直接优化策略。
🛑 可以停在这里
⚡ 走神救援
DQN = Q表换成神经网络。换掉表格两个理由:存不下 + ⭐更根本的是表格没有泛化("敌人在左30像素"和"左31"被当成完全不同的状态)。损失
L=(r+γmaxQ(s',a';θ⁻)−Q(s,a;θ))²—— ⭐把RL变成了监督学习。💀直接这么做会崩,三个问题:样本高度相关、⭐目标在动(追自己尾巴)、数据分布随策略变 —— 合称⭐致命三要素:函数近似+自举+离策略,理论上可发散,第6章的收敛保证在这里失效。三个补丁:①⭐经验回放(打破相关性+经验可重复利用,只有off-policy能用)②⭐目标网络(冻结一份算目标,C步同步一次;太小还在追尾巴、太大目标陈旧)③奖励裁剪+堆叠4帧。⚠️四个代码坑:目标忘了no_grad()训练直接乱、忘了*(1-d)、用MSE不用Huber、目标网络同步频率。Double DQN一行改动:主网络选动作、目标网络评估,治最大化偏差。⭐天花板:只能离散动作(max_a要遍历所有动作),方向盘转多少度/生成哪个token都没法取max → 第9章换路:不估价值,直接优化策略。
下一节 👉 09-策略梯度.md ⭐