Extract shared maze code into common package (Gymnasium-style env)
This commit is contained in:
+128
@@ -0,0 +1,128 @@
|
||||
# DQN 走迷宫:原理与推导
|
||||
|
||||
> 对应代码 `dqn/dqn_maze.py`(PyTorch 实现)。本文不含代码,只讲数学。行内公式用 `$...$`、独立公式用 `$$...$$`,Gitea 可直接渲染。
|
||||
|
||||
---
|
||||
|
||||
## 一、为什么需要 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)的监督——目标会随学习移动。
|
||||
|
||||
---
|
||||
|
||||
## 三、两个关键技巧
|
||||
|
||||
### 3.1 经验回放(Replay Buffer)
|
||||
|
||||
把每条经历 $(s,a,r,s',done)$ 存入一个容量有限的缓冲池,训练时**随机抽一个小批量**更新。
|
||||
|
||||
**为什么必须这么做**: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 步 | 同步频率 |
|
||||
|
||||
---
|
||||
|
||||
## 七、运行与观测
|
||||
|
||||
运行参数:`--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公式推导详解.md》(理论)、《q-learning.md》(表格法原理)。*
|
||||
Reference in New Issue
Block a user