📑 本页目录(点开跳转)
08 · DQN:状态存不下了怎么办
⏱ 30 分钟 | ⭐ 深度强化学习的起点
🎯 一句话
把 Q 表换成神经网络。 听起来只是个小改动,但它会让前面所有的收敛保证全部失效 —— 这一章讲那三个救命的补丁。
💥 一、为什么必须换掉表格
对照
Q 表的大小 = 状态数 × 动作数
格子迷宫 10×10: 100 × 4 = 400 个格子 ✅ 存得下
Atari 一帧画面: 256^(84×84×4) 个状态 💀 宇宙里没这么多原子
而且表格还有一个更根本的问题:
结果对照
🔑 这才是用神经网络的真正理由。 不只是"存不下",是"每个状态都要单独学一遍"根本不现实。
🧠 二、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:样本高度相关
信息关系
问题 2:目标在动(追自己的尾巴)
因果链
问题 3:数据分布随策略改变
信息关系
🔑 这三个合起来叫「致命三要素」(deadly triad): 函数近似 + 自举(bootstrapping)+ 离策略 —— 三者同时出现时,理论上可以发散。 ⭐ 第 6 章那个"表格 Q-learning 保证收敛"的定理,在这里彻底失效了。
🩹 四、三个救命补丁
补丁 1:经验回放(Experience Replay)
结果对照
⭐ 注意这一步只有 off-policy 算法能做 —— 池子里是旧策略采的数据。 🔗 这就是第 6 章说"off-policy 是后面一切基础"的第一次兑现。
补丁 2:目标网络(Target Network)
💡 C 怎么设:常用 1000~10000 步。 太小 → 还是在追尾巴;太大 → 目标太陈旧,学得慢。
补丁 3:奖励裁剪与预处理
结果对照
🔨 五、核心代码
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 的一行改动(值得记住):
# 🧩 骨架:这里只对照两种写法的差别,r / gamma / d / s2 来自你的训练循环
# 原始 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 值,让相似状态之间能够共享信息。
- 经验回放和目标网络缓解相关样本、移动目标造成的不稳定。
- 目标计算、终止掩码和目标网络同步要分别验证,再尝试 Double DQN 等改动。
下一节 👉 09-策略梯度.md ⭐