Files
RL-learning/dqn/dqn-maze.md
T

142 lines
7.0 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# DQN 走迷宫:原理与推导
---
## 一、为什么需要 DQN
Q-Learning 把 Q 值存在一张表里,需要为每个 (状态, 动作) 单独维护一个数。状态空间一大(比如用整张地图做状态),表就装不下、也学不到没见过的格子。
DQN 的思路:**用神经网络把 Q 值"拟合"出来**,而不是逐格记录
$$Q(s,a)\;\approx\;Q(s,a;\theta)$$
其中 $\theta$ 是网络参数。输入状态 $s$,输出各动作的价值。见过 $(s,a)$ 后学到的规律会被网络**泛化**到相近的状态(虽然本迷宫是离散格,但组件完全通用)。
迷宫的问题设定(状态、动作、奖励、转移)与 Q-Learning 完全相同,见 `q-learning.md` 第一、二节;下面直接讲 DQN 自己的部分。
---
## 二、损失函数:把 TD 更新变成监督学习
Q-Learning 的更新是 $Q(s,a)\leftarrow Q(s,a)+\alpha\,\delta$,其中 TD 误差
$$\delta = r+\gamma\max_{a'}Q(s',a')-Q(s,a)$$
DQN 把它改写成**最小化均方误差**的目标:
$$L(\theta)=\mathbb{E}_{(s,a,r,s')}\Big[\big(Y-Q(s,a;\theta)\big)^{2}\Big]$$
其中目标值 $Y$ 用"冻结"的参数 $\theta^{-}$(目标网络)计算:
$$Y=r+\gamma\max_{a'}Q(s',a';\theta^{-})$$
对 $\theta$ 求梯度($Y$ 视为常数,不做梯度):
$$\nabla_{\theta}L=\mathbb{E}\Big[-2\big(Y-Q(s,a;\theta)\big)\,\nabla_{\theta}Q(s,a;\theta)\Big]$$
这正是把贝尔曼最优方程的**样本近似**当成监督学习的"标签"$Y$ 当标签、$Q(s,a;\theta)$ 当预测,用梯度下降让预测逼近标签。$Y$ 本身又依赖旧参数 $\theta^{-}$,所以这是"自举式"bootstrap)的监督——目标会随学习移动。
**符号说明**
| 符号 | 含义 |
|---|---|
| $\theta$ | 主网络参数(全部权重矩阵),每步梯度更新 |
| $\theta^{-}$ | 目标网络参数,$\theta$ 的滞后副本,每 $C$ 步同步一次;算 $Y$ 时冻结、不求梯度 |
| $Q(s,a;\theta)$ | 输入状态 $s$,网络输出的动作 $a$ 的 Q 值("预测" |
| $Y$ | TD 目标("标签"):本步实际奖励 $r$ 加 $\gamma$ 倍下一状态的最优价值估计 |
| $r,\ s',\ a'$ | 本步实际观测到的值:拿到的奖励、到达的新状态、新状态下设想尝试的动作 |
| $\max_{a'}$ | 对 $s'$ 处全部 4 个动作的 Q 值取最大,即"假设下一步走最优" |
| $\gamma$ | 折扣因子,本实现取 0.99 |
| $\mathbb{E}_{(s,a,r,s')}$ | 对经验转移的分布求期望;实现上就是回放池小批量上的平均 |
| $L(\theta)$ | MSE 损失:预测与标签之差的平方的期望 |
| $\nabla_{\theta}L$ | 损失对 $\theta$ 的梯度;Adam 沿其反方向更新参数 |
---
## 三、两个关键技巧
### 3.1 经验回放(Replay Buffer
把每条经历 $(s,a,r,s',terminated)$ 存入一个容量有限的缓冲池,训练时**随机抽一个小批量**更新。注意存的是 `terminated`(到达终点)而非 `truncated`(超时截断):截断只是回合被强制叫停,未来价值依然存在,TD 目标不应因此丢掉 bootstrap 项。
**为什么必须这么做**:TD 目标依赖下一个状态,而连续几步的状态高度相关。若直接用当前轨迹在线更新,梯度在时间上强相关,损失会剧烈震荡甚至发散;随机抽样把这些样本打乱成近似独立同分布,相当于稳定的"数据集",梯度才像普通监督学习那样平稳下降。此外一条经验可被多次复用,样本效率更高。
### 3.2 目标网络(Target Network
若 $Y$ 也用同一个网络 $\theta$ 计算,则"标签"和"预测"同时变化,更新目标一直在追着自己跑(自举放大),训练极易振荡。
解决办法:维护一份**滞后副本** $\theta^{-}$,用 $\theta^{-}$ 算 $Y$,而 $\theta$ 只负责预测并求梯度;每隔固定步数 $C$ 把 $\theta$ 复制给 $\theta^{-}$。这保证了一个更新周期内标签固定,等价于把损失函数稳定下来再优化。
---
## 四、探索与状态表示
- **探索**:与 Q-Learning 相同,用 ε-greedy:以概率 $\varepsilon$ 均匀随机、否则取 $\arg\max_{a}Q(s,a;\theta)$;ε 按 $\varepsilon_k=\max(\varepsilon_{\min},\varepsilon_0\lambda^{k})$ 指数衰减。
- **状态表示**:网络输入需要数字向量。本实现把格子坐标编码成 one-hot 向量(长度 $n\times n$,当前位置为 1),可扩展到"把整张地图作为网格图像 + 卷积"的标准 DQN 形式。
---
## 五、与 Q-Learning 的对比
| | Q-Learning | DQN |
|---|---|---|
| Q 的载体 | 表格 $Q[r,c,a]$ | 神经网络 $Q(s,a;\theta)$ |
| 更新目标 | $r+\gamma\max_{a'}Q(s',a')$(查表)| 同上,但用目标网络 $\theta^{-}$ |
| 优化手段 | 直接加法修正 | 对 MSE 损失做梯度下降 |
| 额外组件 | 无 | 经验回放、目标网络 |
| 泛化 | 无(没见过的格子没值)| 有(近似相近状态)|
| 收敛性 | 有理论保证(Watkins)| **无严格保证**,靠工程技巧稳定 |
两者本质相同:都在做贝尔曼最优方程的随机近似,只是函数逼近方式不同。
---
## 六、超参数(默认值)
| 超参数 | 默认 | 说明 |
|---|---|---|
| 网络宽度 hidden | 128 | 两层隐层 |
| 学习率 lr | 1e-3 | Adam 优化器 |
| 批量大小 batch | 64 | 每次从回放池抽样数 |
| 回放容量 | 20000 | 过大滞后、过小样本相关 |
| γ | 0.99 | 折扣因子 |
| ε 衰减 | 1.0 → 0.02 | 每局 ×0.995 |
| 目标更新 C | 200 步 | 同步频率 |
---
## 七、运行与观测
在仓库根目录运行:`python -m dqn.dqn_maze --plot`;参数:`--rows/--cols`(默认 9,须奇数)、`--seed``--episodes`(默认 800)、`--braid``--hidden``--plot`。输出:迷宫、策略箭头、贪心路径、成功率统计;`--plot` 生成两张图。
| 图 | 判读 |
|---|---|
| 逐局回报 | 总体上升(噪声大,看滑动平均)|
| 滑动成功率 | 趋近 1 |
| 单局步数 vs BFS 红线 | 降到红线 = 学会最短路 |
| TD 损失(对数轴)| 下降即 $Q(s,a;\theta)$ 逼近 $Q^{*}$;平台或回升提示超参/技巧问题 |
| V(s) 快照 | 价值从终点回传的过程 |
常见问题:
| 现象 | 原因 | 调整 |
|---|---|---|
| 损失震荡不降 | 无目标网络或回放太小 | 检查目标网络、增大回放容量 |
| 学不到终点 | 探索不足或 γ 太小 | 提高 ε 下限、γ 取 0.99 |
| 前期好后期崩 | 回放池被旧数据稀释、学习率过大 | 降 lr、增容量 |
| 结果不稳定 | DQN 本身随机性强 | 固定 seed、多看几次 |
---
## 八、小结
DQN = Q-Learning 的公式 + 神经网络的表达力 + 两条工程稳定技巧。核心仍是那条损失
$$L(\theta)=\mathbb{E}\big[(r+\gamma\max_{a'}Q(s',a';\theta^{-})-Q(s,a;\theta))^{2}\big]$$
吃透这条式子,再理解回放和目标网络各自解决的"相关性与自举"两个病,DQN 的骨架就齐了——换成连续动作空间(DDPG)或图像输入(CNN)只是改编码器和输出头。
---
*配套阅读:`q_learning/q-learning.md`(表格法原理与推导)。*