🏠 总目录📚 本教程 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 的一行改动(值得记住):

# 原始 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 章换路:不估价值,直接优化策略

🛑 可以停在这里

走神救援

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

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