🏠 总目录📚 本教程 08 · DQN ← →
📑 本页目录(点开跳转)

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)$$

输入:状态 s比如 4 帧游戏画面
→
CNN参数 θ 的 Q 值近似器
→
输出:每个动作的 Q 值Q(s,←)、Q(s,→)、Q(s,↑)、Q(s,↓) 一次全出来 ⭐

💡 一个小设计但很重要:输出所有动作的 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 步才同步一次】 训练网络 θ Q(s,a; θ) 每一步都被梯度更新 目标网络 θ⁻ Q(s',a'; θ⁻) 参数冻住不动 每隔 C 步拷贝参数 L = (r + γ max Q(s',a'; θ⁻ ) − Q(s,a; θ))² 冻结的那份 ✅ 目标在 C 步内是【固定】的 → 变回了标准的监督学习 ✅ 打破了“追自己尾巴”的循环
两条线:训练网络每一步都在动,目标网络每隔 C 步才拷一次参数 —— 被冻住的那个 θ⁻ 就是「追自己尾巴」这个循环被切断的地方

💡 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 的一行改动(值得记住):

# 🧩 骨架:这里只对照两种写法的差别,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 里裁的是同一个东西

✅ 检查点

  1. 换掉表格的两个理由是什么?哪个更根本?
  2. DQN 的损失函数是什么?它把 RL 变成了什么问题?
  3. 「致命三要素」是哪三个?它意味着什么?
  4. 经验回放解决什么问题?为什么只有 off-policy 能用?
  5. 目标网络解决什么问题?C 太小太大分别会怎样?
  6. 代码里忘了 torch.no_grad() 会怎样?忘了 *(1-d) 呢?
  7. Double DQN 的一行改动是什么?解决什么?
  8. DQN 的最大天花板是什么?为什么它逼出了第 9 章?
👀 答案
  1. ①状态太多存不下(Atari 是 256^(84×84×4))②表格没有泛化能力——"敌人在左30像素"和"左31像素"被当成完全不同的状态。第二个更根本:不只是存不下,是每个状态都要单独学一遍不现实。
  2. L = (r + γ max Q(s',a';θ⁻) − Q(s,a;θ))²。它把 RL 变成了监督学习问题(拟合一个回归目标),于是能用梯度下降。
  3. 函数近似 + 自举 + 离策略。三者同时出现时理论上可以发散——第 6 章的表格收敛保证在这里彻底失效。
  4. 解决样本高度相关(连续帧几乎一样,梯度方向高度相关)。只有 off-policy 能用因为池子里是旧策略采的数据。
  5. 解决目标在动、追自己尾巴的问题。C 太小还是在追尾巴;C 太大目标太陈旧学得慢(常用 1000~10000)。
  6. 忘 no_grad() → 梯度回传到目标网络,训练直接乱掉;忘 *(1-d) → 终止态算进了不存在的未来价值。
  7. 主网络选动作、目标网络评估:a2=q(s2).argmax(); tgt=r+γ·q_target(s2).gather(1,a2)。解决最大化偏差(max 系统性挑中被高估的动作)。
  8. 只能处理离散动作——max_a 要遍历所有动作,动作连续就没法算。所以第 9 章换路:不估价值,直接优化策略。

🛑 可以停在这里

⚡ 走神救援

先记住这几件事

下一节 👉 09-策略梯度.md ⭐

打卡记录保存在你的浏览器里,首页能看到总进度