Extract shared maze code into common package (Gymnasium-style env)
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
*
|
||||
!.gitignore
|
||||
!common/
|
||||
!q_learning/
|
||||
!dqn/
|
||||
!*.md
|
||||
!*.py
|
||||
|
||||
+303
@@ -0,0 +1,303 @@
|
||||
"""Q-Learning 与 DQN 共用的迷宫环境、展示与绘图工具。"""
|
||||
from collections import deque
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
START, END, WALL, ROAD = 0, 1, 2, 3 # 分别是起点 终点 障碍 道路
|
||||
ACTION_DELTAS = ((-1, 0), (1, 0), (0, -1), (0, 1)) # 动作: 0上 1下 2左 3右
|
||||
ARROWS = "↑↓←→"
|
||||
CHARS = {START: "S", END: "G", WALL: "#", ROAD: "."}
|
||||
|
||||
|
||||
def generate_maze(rows, cols, seed=0, braid=0.0):
|
||||
"""按 seed 随机生成可解迷宫(递归回溯法),行列须为不小于 3 的奇数。
|
||||
|
||||
保证起点 (0,0) 到终点 (rows-1, cols-1) 连通。
|
||||
braid: 0~1, 生成后按该概率拆掉剩余内墙形成环路。
|
||||
0 = 完美迷宫(唯一通路, 长走廊多);越大环路越多、可选路径越多。
|
||||
"""
|
||||
if rows < 3 or cols < 3 or rows % 2 == 0 or cols % 2 == 0:
|
||||
raise ValueError("行列必须为不小于 3 的奇数")
|
||||
rng = np.random.default_rng(seed)
|
||||
maze = np.full((rows, cols), WALL, dtype=np.int8)
|
||||
maze[0, 0] = ROAD
|
||||
stack = [(0, 0)]
|
||||
while stack:
|
||||
r, c = stack[-1]
|
||||
neighbors = [
|
||||
(r + dr, c + dc, r + dr // 2, c + dc // 2)
|
||||
for dr, dc in ((-2, 0), (2, 0), (0, -2), (0, 2))
|
||||
if 0 <= r + dr < rows and 0 <= c + dc < cols
|
||||
and maze[r + dr, c + dc] == WALL
|
||||
]
|
||||
if neighbors:
|
||||
nr, nc, mr, mc = neighbors[int(rng.integers(len(neighbors)))]
|
||||
maze[mr, mc] = maze[nr, nc] = ROAD
|
||||
stack.append((nr, nc))
|
||||
else:
|
||||
stack.pop()
|
||||
if braid > 0:
|
||||
for r in range(1, rows - 1):
|
||||
for c in range(1, cols - 1):
|
||||
if maze[r, c] == WALL and rng.random() < braid and (
|
||||
(maze[r - 1, c] != WALL and maze[r + 1, c] != WALL)
|
||||
or (maze[r, c - 1] != WALL and maze[r, c + 1] != WALL)
|
||||
):
|
||||
maze[r, c] = ROAD # 拆掉两侧都是路的墙, 形成环路
|
||||
maze[0, 0], maze[rows - 1, cols - 1] = START, END
|
||||
return maze
|
||||
|
||||
|
||||
class MazeEnv:
|
||||
"""Gymnasium 风格的迷宫环境(不依赖 gymnasium, 仅遵循其 API 约定)。
|
||||
|
||||
状态 obs 为格子坐标 (r, c);动作: 0上 1下 2左 3右。
|
||||
奖励: 每步 step_reward;撞墙/越界 wall_penalty(原地不动);
|
||||
到达终点 goal_reward 且 terminated。
|
||||
step() 返回 (obs, reward, terminated, truncated, info):
|
||||
terminated = 到达终点(真正的终止状态);
|
||||
truncated = 达到 max_steps 步数上限(超时截断, 非终止状态)。
|
||||
"""
|
||||
|
||||
def __init__(self, maze, start_pos=None,
|
||||
step_reward=-1.0, wall_penalty=-10.0, goal_reward=100.0, max_steps=200):
|
||||
self.maze = np.asarray(maze)
|
||||
self.shape = self.maze.shape
|
||||
self.n_actions = len(ACTION_DELTAS)
|
||||
self.start_pos = self._find(START) if start_pos is None else tuple(start_pos)
|
||||
self.goal = self._find(END)
|
||||
|
||||
r, c = self.start_pos
|
||||
if not (0 <= r < self.shape[0] and 0 <= c < self.shape[1]) or self.maze[r, c] == WALL:
|
||||
raise ValueError(f"非法起点: {self.start_pos}")
|
||||
if self.start_pos == self.goal:
|
||||
raise ValueError("起点不能与终点重合")
|
||||
|
||||
self.step_reward = step_reward
|
||||
self.wall_penalty = wall_penalty
|
||||
self.goal_reward = goal_reward
|
||||
self.max_steps = max_steps
|
||||
self.steps = 0
|
||||
self._pos = self.start_pos
|
||||
|
||||
def _find(self, cell):
|
||||
hits = np.argwhere(self.maze == cell)
|
||||
if len(hits) == 0:
|
||||
raise ValueError(f"迷宫中找不到格子类型 {cell}")
|
||||
return tuple(int(v) for v in hits[0])
|
||||
|
||||
def reset(self, *, seed=None, options=None):
|
||||
"""重置到起点, 返回 (obs, info)。seed/options 仅为 API 兼容保留。"""
|
||||
self._pos = self.start_pos
|
||||
self.steps = 0
|
||||
return self._pos, {}
|
||||
|
||||
def step(self, action):
|
||||
"""执行动作, 返回 (obs, reward, terminated, truncated, info)。"""
|
||||
if not 0 <= action < self.n_actions:
|
||||
raise ValueError(f"非法动作: {action}")
|
||||
|
||||
dr, dc = ACTION_DELTAS[action]
|
||||
nr, nc = self._pos[0] + dr, self._pos[1] + dc
|
||||
self.steps += 1
|
||||
|
||||
terminated = False
|
||||
if (
|
||||
not (0 <= nr < self.shape[0] and 0 <= nc < self.shape[1])
|
||||
or self.maze[nr][nc] == WALL
|
||||
):
|
||||
# 遇到边界或者撞墙: 原地不动
|
||||
reward = self.wall_penalty
|
||||
else:
|
||||
self._pos = (nr, nc)
|
||||
terminated = self._pos == self.goal
|
||||
reward = self.goal_reward if terminated else self.step_reward
|
||||
|
||||
truncated = not terminated and self.steps >= self.max_steps
|
||||
return self._pos, reward, terminated, truncated, {}
|
||||
|
||||
|
||||
def shortest_path_len(env):
|
||||
"""BFS 求起点到终点的最短步数(不通返回 -1)。"""
|
||||
dist = {env.start_pos: 0}
|
||||
queue = deque([env.start_pos])
|
||||
while queue:
|
||||
pos = queue.popleft()
|
||||
if pos == env.goal:
|
||||
return dist[pos]
|
||||
for dr, dc in ACTION_DELTAS:
|
||||
nr, nc = pos[0] + dr, pos[1] + dc
|
||||
nxt = (nr, nc)
|
||||
if (0 <= nr < env.shape[0] and 0 <= nc < env.shape[1]
|
||||
and env.maze[nr, nc] != WALL and nxt not in dist):
|
||||
dist[nxt] = dist[pos] + 1
|
||||
queue.append(nxt)
|
||||
return -1
|
||||
|
||||
|
||||
def evaluate(env, act, episodes=100):
|
||||
"""纯贪心策略(无探索)评测, 返回 (成功率, 成功局平均步数)。
|
||||
|
||||
act(pos) -> 动作编号, 由各算法提供(查表或前向推理)。
|
||||
"""
|
||||
wins, steps_list = 0, []
|
||||
for _ in range(episodes):
|
||||
obs, _ = env.reset()
|
||||
terminated = truncated = False
|
||||
while not (terminated or truncated):
|
||||
obs, _, terminated, truncated, _ = env.step(int(act(obs)))
|
||||
if terminated:
|
||||
wins += 1
|
||||
steps_list.append(env.steps)
|
||||
avg_steps = float(np.mean(steps_list)) if steps_list else float("nan")
|
||||
return wins / episodes, avg_steps
|
||||
|
||||
|
||||
def show_maze(env):
|
||||
print(f"迷宫 {env.shape[0]}x{env.shape[1]}(S起点 G终点 #墙):")
|
||||
for r in range(env.shape[0]):
|
||||
print("".join(CHARS[env.maze[r][c]] for c in range(env.shape[1])))
|
||||
|
||||
|
||||
def show_policy(env, act):
|
||||
print("学到的策略(每格最优动作):")
|
||||
for r in range(env.shape[0]):
|
||||
row = ""
|
||||
for c in range(env.shape[1]):
|
||||
if env.maze[r][c] == WALL:
|
||||
row += " # "
|
||||
elif (r, c) == env.goal:
|
||||
row += " G "
|
||||
else:
|
||||
row += f" {ARROWS[int(act((r, c)))]} "
|
||||
print(row)
|
||||
|
||||
|
||||
def show_path(env, act):
|
||||
print("\n贪心策略走出的一条路径:")
|
||||
obs, path, terminated = env.reset()[0], [env.start_pos], False
|
||||
truncated = False
|
||||
while not (terminated or truncated):
|
||||
obs, _, terminated, truncated, _ = env.step(int(act(obs)))
|
||||
path.append(obs)
|
||||
grid = [[CHARS[env.maze[r][c]] for c in range(env.shape[1])]
|
||||
for r in range(env.shape[0])]
|
||||
for i, (r, c) in enumerate(path):
|
||||
if (r, c) not in (env.start_pos, env.goal):
|
||||
grid[r][c] = str(i % 10) # 若重复经过同一格,显示最后一步序号
|
||||
print("\n".join(" ".join(row) for row in grid))
|
||||
print(f"到达终点: {'是' if terminated else '否'},共 {len(path) - 1} 步")
|
||||
|
||||
|
||||
# ─────────────────── 训练过程可视化(需 matplotlib) ───────────────────
|
||||
def _setup_plt():
|
||||
import matplotlib
|
||||
|
||||
matplotlib.use("Agg") # 不弹窗,直接保存图片
|
||||
import matplotlib.pyplot as plt
|
||||
plt.rcParams["font.sans-serif"] = ["Microsoft YaHei", "SimHei"] # 中文字体
|
||||
plt.rcParams["axes.unicode_minus"] = False # 正常显示负号
|
||||
return plt
|
||||
|
||||
|
||||
def moving_avg(x, k=20):
|
||||
x = np.asarray(x, dtype=float)
|
||||
if len(x) < k:
|
||||
return np.array([])
|
||||
return np.convolve(x, np.ones(k) / k, mode="valid")
|
||||
|
||||
|
||||
def _plot_ma(ax, x, window, color):
|
||||
x = np.asarray(x, dtype=float)
|
||||
ax.plot(x, color="gray", alpha=0.35, label="逐局值")
|
||||
ma = moving_avg(x, window)
|
||||
if len(ma):
|
||||
ax.plot(np.arange(window - 1, window - 1 + len(ma)), ma,
|
||||
color=color, lw=2, label=f"滑动平均(窗口{window})")
|
||||
ax.set_xlabel("训练局数 Episode")
|
||||
ax.legend()
|
||||
|
||||
|
||||
def plot_training(env, rewards, successes, ep_steps, extra,
|
||||
window=20, out="training.png", title="走迷宫训练过程"):
|
||||
"""四联图: 回报曲线 / 滑动成功率 / 单局步数 / 第4格由 extra 提供。
|
||||
|
||||
extra: (数值序列, y轴标签, 子图标题, 是否对数y轴)。
|
||||
"""
|
||||
plt = _setup_plt()
|
||||
fig, axes = plt.subplots(2, 2, figsize=(12, 8))
|
||||
fig.suptitle(title, fontsize=15)
|
||||
|
||||
ax = axes[0][0]
|
||||
_plot_ma(ax, rewards, window, "C0")
|
||||
ax.set_ylabel("单局回报")
|
||||
ax.set_title("逐局回报(越高越好)")
|
||||
|
||||
ax = axes[0][1]
|
||||
_plot_ma(ax, successes, window, "C1")
|
||||
ax.set_ylim(-0.05, 1.05)
|
||||
ax.set_ylabel("成功率")
|
||||
ax.set_title(f"滑动成功率(窗口{window})")
|
||||
|
||||
ax = axes[1][0]
|
||||
_plot_ma(ax, ep_steps, window, "C2")
|
||||
sp = shortest_path_len(env)
|
||||
if sp >= 0:
|
||||
ax.axhline(sp, color="red", ls="--", lw=1.5, label=f"最短路 {sp} 步")
|
||||
ax.set_ylabel("单局步数")
|
||||
ax.set_title("单局步数(降到红线 = 学会最短路)")
|
||||
ax.legend()
|
||||
|
||||
values, ylabel, subtitle, log_scale = extra
|
||||
ax = axes[1][1]
|
||||
ax.plot(np.asarray(values, dtype=float), color="C3", lw=1.2 if log_scale else 2)
|
||||
if log_scale:
|
||||
ax.set_yscale("log")
|
||||
ax.set_xlabel("训练局数 Episode")
|
||||
ax.set_ylabel(ylabel)
|
||||
ax.set_title(subtitle)
|
||||
|
||||
fig.tight_layout(rect=[0, 0, 1, 0.96])
|
||||
fig.savefig(out, dpi=150)
|
||||
print(f"\n图片已保存: {out}")
|
||||
|
||||
|
||||
def plot_snapshots(env, snapshots, v_final, out="snapshots.png",
|
||||
title="V(s)=max Q 的学习过程(价值从 G 回传)"):
|
||||
"""各训练阶段的价值热力图。
|
||||
|
||||
snapshots/v_final: (标签, V二维数组),V 为该阶段每格的 max Q,墙格可为 nan。
|
||||
"""
|
||||
plt = _setup_plt()
|
||||
panels = list(snapshots) + [("最终", v_final)]
|
||||
vals = [np.asarray(V, dtype=float)[~np.isnan(V)] for _, V in panels]
|
||||
vmin = min(v.min() for v in vals)
|
||||
vmax = max(v.max() for v in vals)
|
||||
if vmin == vmax:
|
||||
vmin, vmax = vmin - 1, vmax + 1
|
||||
|
||||
n = len(panels)
|
||||
fig, axes = plt.subplots(1, n, figsize=(3 * n, 3.4))
|
||||
axes = np.atleast_1d(axes)
|
||||
masked_wall = env.maze == WALL
|
||||
for ax, (tag, V) in zip(axes, panels):
|
||||
V = np.asarray(V, dtype=float)
|
||||
Vm = np.ma.masked_where(masked_wall, V)
|
||||
im = ax.imshow(Vm, cmap="viridis", vmin=vmin, vmax=vmax, origin="upper")
|
||||
fs = 8 if max(env.shape) <= 7 else 5
|
||||
for r in range(env.shape[0]):
|
||||
for c in range(env.shape[1]):
|
||||
if masked_wall[r, c]:
|
||||
ax.text(c, r, "█", ha="center", va="center",
|
||||
color="black", fontsize=fs)
|
||||
else:
|
||||
ax.text(c, r, f"{V[r, c]:.0f}", ha="center", va="center",
|
||||
color="white", fontsize=fs)
|
||||
ax.set_xticks([])
|
||||
ax.set_yticks([])
|
||||
ax.set_title(f"第 {tag} 局后" if tag != "最终" else "训练结束")
|
||||
fig.suptitle(title, fontsize=13)
|
||||
fig.colorbar(im, ax=axes, fraction=0.03, pad=0.02)
|
||||
fig.savefig(out, dpi=150, bbox_inches="tight")
|
||||
print(f"图片已保存: {out}")
|
||||
+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》(表格法原理)。*
|
||||
+202
@@ -0,0 +1,202 @@
|
||||
import argparse
|
||||
from collections import deque
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from common.maze import (
|
||||
WALL,
|
||||
MazeEnv,
|
||||
evaluate,
|
||||
generate_maze,
|
||||
plot_snapshots,
|
||||
plot_training,
|
||||
show_maze,
|
||||
show_path,
|
||||
show_policy,
|
||||
)
|
||||
|
||||
|
||||
class DQN(nn.Module):
|
||||
"""输入: 状态编码(one-hot 位置向量); 输出: 4 个动作的 Q 值。"""
|
||||
|
||||
def __init__(self, n_states, n_actions, hidden=128):
|
||||
super().__init__()
|
||||
self.net = nn.Sequential(
|
||||
nn.Linear(n_states, hidden),
|
||||
nn.ReLU(),
|
||||
nn.Linear(hidden, hidden),
|
||||
nn.ReLU(),
|
||||
nn.Linear(hidden, n_actions),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
|
||||
class ReplayBuffer:
|
||||
"""经验回放池: 存储 (s, a, r, s', done), 随机抽样打破时序相关性。"""
|
||||
|
||||
def __init__(self, capacity, seed=0):
|
||||
self.buf = deque(maxlen=capacity)
|
||||
self.rng = np.random.default_rng(seed)
|
||||
|
||||
def push(self, *transition):
|
||||
self.buf.append(transition)
|
||||
|
||||
def sample(self, batch_size):
|
||||
idx = self.rng.choice(len(self.buf), batch_size, replace=False)
|
||||
batch = [self.buf[i] for i in idx]
|
||||
return tuple(np.stack(col) for col in zip(*batch))
|
||||
|
||||
def __len__(self):
|
||||
return len(self.buf)
|
||||
|
||||
|
||||
def encode_state(pos, n_states, n_cols):
|
||||
"""把格子坐标 (r, c) 编码成 one-hot 向量。"""
|
||||
one = np.zeros(n_states, dtype=np.float32)
|
||||
one[pos[0] * n_cols + pos[1]] = 1.0
|
||||
return one
|
||||
|
||||
|
||||
def optimize(net, target, opt, buffer, batch_size, gamma):
|
||||
"""一次梯度步: 最小化 (r + γ·max Q(s',·;θ⁻) − Q(s,a;θ))²"""
|
||||
s, a, r, s2, done = buffer.sample(batch_size)
|
||||
s = torch.as_tensor(s)
|
||||
a = torch.as_tensor(a, dtype=torch.long).unsqueeze(1)
|
||||
r = torch.as_tensor(r, dtype=torch.float32)
|
||||
s2 = torch.as_tensor(s2, dtype=torch.float32)
|
||||
done = torch.as_tensor(done, dtype=torch.float32)
|
||||
|
||||
q = net(s).gather(1, a).squeeze(1)
|
||||
with torch.no_grad():
|
||||
y = r + gamma * target(s2).max(1).values * (1.0 - done)
|
||||
loss = F.mse_loss(q, y)
|
||||
|
||||
opt.zero_grad()
|
||||
loss.backward()
|
||||
opt.step()
|
||||
return float(loss.item())
|
||||
|
||||
|
||||
def value_map(env, net, encode):
|
||||
"""计算每个格子的 V(s)=max_a Q(s,a), 墙为 nan。"""
|
||||
rows, cols = env.shape
|
||||
V = np.full((rows, cols), np.nan)
|
||||
with torch.no_grad():
|
||||
for r in range(rows):
|
||||
for c in range(cols):
|
||||
if env.maze[r, c] == WALL:
|
||||
continue
|
||||
x = torch.as_tensor(encode((r, c))).unsqueeze(0)
|
||||
V[r, c] = float(net(x).max().item())
|
||||
return V
|
||||
|
||||
|
||||
def train(env, episodes=800, hidden=128, batch_size=64, buffer_capacity=20000,
|
||||
gamma=0.99, lr=1e-3, epsilon=1.0, eps_min=0.02, eps_decay=0.995,
|
||||
target_update=200, seed=0, snap_at=()):
|
||||
torch.manual_seed(seed)
|
||||
np.random.seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
n_states = env.shape[0] * env.shape[1]
|
||||
net = DQN(n_states, env.n_actions, hidden).to(device)
|
||||
target = DQN(n_states, env.n_actions, hidden).to(device)
|
||||
target.load_state_dict(net.state_dict())
|
||||
opt = torch.optim.Adam(net.parameters(), lr=lr)
|
||||
buffer = ReplayBuffer(buffer_capacity, seed=seed)
|
||||
rng = np.random.default_rng(seed + 1)
|
||||
|
||||
def encode(pos):
|
||||
return encode_state(pos, n_states, env.shape[1])
|
||||
|
||||
losses, rewards, successes, ep_steps, eps_hist = [], [], [], [], []
|
||||
snapshots = []
|
||||
snap_at = sorted({int(e) for e in snap_at if 0 < int(e) <= episodes})
|
||||
si = 0
|
||||
total_steps = 0
|
||||
|
||||
for ep in range(1, episodes + 1):
|
||||
pos, _ = env.reset()
|
||||
total_r, ep_loss, n_loss = 0.0, 0.0, 0
|
||||
terminated = truncated = False
|
||||
|
||||
while not (terminated or truncated):
|
||||
# ε-greedy
|
||||
if rng.random() < epsilon:
|
||||
a = int(rng.integers(env.n_actions))
|
||||
else:
|
||||
with torch.no_grad():
|
||||
x = torch.as_tensor(encode(pos)).unsqueeze(0).to(device)
|
||||
a = int(net(x).argmax().item())
|
||||
|
||||
nxt, r, terminated, truncated, _ = env.step(a)
|
||||
# 超时截断(truncated)不代表环境终止, TD 目标不该被截断, 故只存 terminated
|
||||
buffer.push(encode(pos), a, r, encode(nxt), terminated)
|
||||
|
||||
if len(buffer) >= batch_size:
|
||||
ep_loss += optimize(net, target, opt, buffer, batch_size, gamma)
|
||||
n_loss += 1
|
||||
total_steps += 1
|
||||
if total_steps % target_update == 0:
|
||||
target.load_state_dict(net.state_dict())
|
||||
|
||||
pos, total_r = nxt, total_r + r
|
||||
|
||||
epsilon = max(eps_min, epsilon * eps_decay)
|
||||
losses.append(ep_loss / n_loss if n_loss else float("nan"))
|
||||
rewards.append(total_r)
|
||||
successes.append(terminated)
|
||||
ep_steps.append(env.steps)
|
||||
eps_hist.append(epsilon)
|
||||
if si < len(snap_at) and ep == snap_at[si]:
|
||||
snapshots.append((ep, value_map(env, net, encode)))
|
||||
si += 1
|
||||
|
||||
return net, device, encode, (losses, rewards, successes, ep_steps, eps_hist, snapshots)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="DQN 走迷宫(PyTorch 深度强化学习实现)")
|
||||
parser.add_argument("--seed", type=int, default=0, help="随机种子")
|
||||
parser.add_argument("--rows", type=int, default=9, help="迷宫行数(奇数)")
|
||||
parser.add_argument("--cols", type=int, default=9, help="迷宫列数(奇数)")
|
||||
parser.add_argument("--episodes", type=int, default=800, help="训练局数")
|
||||
parser.add_argument("--braid", type=float, default=0.0, help="拆墙成环概率 0~1")
|
||||
parser.add_argument("--hidden", type=int, default=128, help="隐藏层宽度")
|
||||
parser.add_argument("--plot", action="store_true",
|
||||
help="生成 dqn_training.png 与 dqn_snapshots.png")
|
||||
args = parser.parse_args()
|
||||
|
||||
env = MazeEnv(generate_maze(args.rows, args.cols, seed=args.seed, braid=args.braid))
|
||||
show_maze(env)
|
||||
snap_at = (1, 20, 100, 300, 500) if args.plot else ()
|
||||
net, device, encode, hist = train(
|
||||
env, episodes=args.episodes, hidden=args.hidden, seed=args.seed, snap_at=snap_at)
|
||||
losses, rewards, successes, ep_steps, eps_hist, snapshots = hist
|
||||
|
||||
def act(pos):
|
||||
with torch.no_grad():
|
||||
x = torch.as_tensor(encode(pos)).unsqueeze(0).to(device)
|
||||
return int(net(x).argmax().item())
|
||||
|
||||
show_policy(env, act)
|
||||
show_path(env, act)
|
||||
rate, avg_steps = evaluate(env, act)
|
||||
print(f"\n前50局平均回报: {np.mean(rewards[:50]):6.1f} "
|
||||
f"后50局平均回报: {np.mean(rewards[-50:]):6.1f}")
|
||||
print(f"后50局成功率: {np.mean(successes[-50:]):.0%}")
|
||||
print(f"评测(100局贪心): 成功率 {rate:.0%},平均 {avg_steps:.1f} 步到达终点")
|
||||
if args.plot:
|
||||
plot_training(env, rewards, successes, ep_steps,
|
||||
extra=(losses, "TD 损失 (MSE, log)", "TD 损失(下降即逼近 Q*)", True),
|
||||
out="dqn_training.png", title="DQN 走迷宫训练过程")
|
||||
plot_snapshots(env, snapshots, value_map(env, net, encode),
|
||||
out="dqn_snapshots.png",
|
||||
title="DQN 学到的 V(s)=max Q 随训练的变化(价值从 G 回传)")
|
||||
+30
-284
@@ -1,129 +1,17 @@
|
||||
import argparse
|
||||
from collections import deque
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
START, END, WALL, ROAD = 0, 1, 2, 3 # 分别是起点 终点 障碍 道路
|
||||
ACTION_DELTAS = ((-1, 0), (1, 0), (0, -1), (0, 1)) # 动作: 0上 1下 2左 3右
|
||||
ARROWS = "↑↓←→"
|
||||
CHARS = {START: "S", END: "G", WALL: "#", ROAD: "."}
|
||||
|
||||
|
||||
def generate_maze(rows, cols, seed=0, braid=0.0):
|
||||
"""按 seed 随机生成可解迷宫(递归回溯法),行列须为不小于 3 的奇数。
|
||||
|
||||
保证起点 (0,0) 到终点 (rows-1, cols-1) 连通。
|
||||
braid: 0~1, 生成后按该概率拆掉剩余内墙形成环路。
|
||||
0 = 完美迷宫(唯一通路, 长走廊多);越大环路越多、可选路径越多。
|
||||
"""
|
||||
if rows < 3 or cols < 3 or rows % 2 == 0 or cols % 2 == 0:
|
||||
raise ValueError("行列必须为不小于 3 的奇数")
|
||||
rng = np.random.default_rng(seed)
|
||||
maze = np.full((rows, cols), WALL, dtype=np.int8)
|
||||
maze[0, 0] = ROAD
|
||||
stack = [(0, 0)]
|
||||
while stack:
|
||||
r, c = stack[-1]
|
||||
neighbors = [
|
||||
(r + dr, c + dc, r + dr // 2, c + dc // 2)
|
||||
for dr, dc in ((-2, 0), (2, 0), (0, -2), (0, 2))
|
||||
if 0 <= r + dr < rows and 0 <= c + dc < cols
|
||||
and maze[r + dr, c + dc] == WALL
|
||||
]
|
||||
if neighbors:
|
||||
nr, nc, mr, mc = neighbors[int(rng.integers(len(neighbors)))]
|
||||
maze[mr, mc] = maze[nr, nc] = ROAD
|
||||
stack.append((nr, nc))
|
||||
else:
|
||||
stack.pop()
|
||||
if braid > 0:
|
||||
for r in range(1, rows - 1):
|
||||
for c in range(1, cols - 1):
|
||||
if maze[r, c] == WALL and rng.random() < braid and (
|
||||
(maze[r - 1, c] != WALL and maze[r + 1, c] != WALL)
|
||||
or (maze[r, c - 1] != WALL and maze[r, c + 1] != WALL)
|
||||
):
|
||||
maze[r, c] = ROAD # 拆掉两侧都是路的墙, 形成环路
|
||||
maze[0, 0], maze[rows - 1, cols - 1] = START, END
|
||||
return maze
|
||||
|
||||
|
||||
class MazeEnv:
|
||||
|
||||
def __init__(self, maze, start_pos=None,
|
||||
step_reward=-1.0, wall_penalty=-10.0, goal_reward=100.0, max_steps=200):
|
||||
self.maze = np.asarray(maze)
|
||||
self.shape = self.maze.shape
|
||||
self.n_actions = len(ACTION_DELTAS)
|
||||
self.start_pos = self._find(START) if start_pos is None else tuple(start_pos)
|
||||
self.goal = self._find(END)
|
||||
|
||||
r, c = self.start_pos
|
||||
if not (0 <= r < self.shape[0] and 0 <= c < self.shape[1]) or self.maze[r, c] == WALL:
|
||||
raise ValueError(f"非法起点: {self.start_pos}")
|
||||
if self.start_pos == self.goal:
|
||||
raise ValueError("起点不能与终点重合")
|
||||
|
||||
self.step_reward = step_reward
|
||||
self.wall_penalty = wall_penalty
|
||||
self.goal_reward = goal_reward
|
||||
self.max_steps = max_steps
|
||||
self.steps = 0
|
||||
self._pos = self.start_pos
|
||||
|
||||
def _find(self, cell):
|
||||
hits = np.argwhere(self.maze == cell)
|
||||
if len(hits) == 0:
|
||||
raise ValueError(f"迷宫中找不到格子类型 {cell}")
|
||||
return tuple(int(v) for v in hits[0])
|
||||
|
||||
def reset(self):
|
||||
self._pos = self.start_pos
|
||||
self.steps = 0
|
||||
return self._pos
|
||||
|
||||
def step(self, action):
|
||||
if not 0 <= action < self.n_actions:
|
||||
raise ValueError(f"非法动作: {action}")
|
||||
|
||||
dr, dc = ACTION_DELTAS[action]
|
||||
nr, nc = self._pos[0] + dr, self._pos[1] + dc
|
||||
self.steps += 1
|
||||
|
||||
if (
|
||||
not (0 <= nr < self.shape[0] and 0 <= nc < self.shape[1])
|
||||
or self.maze[nr][nc] == WALL
|
||||
):
|
||||
# 遇到边界或者撞墙
|
||||
return self._pos, self.wall_penalty, False
|
||||
|
||||
self._pos = (nr, nc)
|
||||
|
||||
if self._pos == self.goal:
|
||||
# 到达迷宫终点
|
||||
return self._pos, self.goal_reward, True
|
||||
|
||||
# 每步 -1, 促使走最短路
|
||||
return self._pos, self.step_reward, False
|
||||
|
||||
|
||||
def shortest_path_len(env):
|
||||
"""BFS 求起点到终点的最短步数(不通返回 -1)。"""
|
||||
dist = {env.start_pos: 0}
|
||||
queue = deque([env.start_pos])
|
||||
while queue:
|
||||
pos = queue.popleft()
|
||||
if pos == env.goal:
|
||||
return dist[pos]
|
||||
for dr, dc in ACTION_DELTAS:
|
||||
nr, nc = pos[0] + dr, pos[1] + dc
|
||||
nxt = (nr, nc)
|
||||
if (0 <= nr < env.shape[0] and 0 <= nc < env.shape[1]
|
||||
and env.maze[nr, nc] != WALL and nxt not in dist):
|
||||
dist[nxt] = dist[pos] + 1
|
||||
queue.append(nxt)
|
||||
return -1
|
||||
from common.maze import (
|
||||
MazeEnv,
|
||||
evaluate,
|
||||
generate_maze,
|
||||
plot_snapshots,
|
||||
plot_training,
|
||||
show_maze,
|
||||
show_path,
|
||||
show_policy,
|
||||
)
|
||||
|
||||
|
||||
def train(env, episodes=500, alpha=0.1,
|
||||
@@ -136,8 +24,9 @@ def train(env, episodes=500, alpha=0.1,
|
||||
si = 0
|
||||
|
||||
for ep in range(1, episodes + 1):
|
||||
pos, total_r, done = env.reset(), 0.0, False
|
||||
while env.steps < env.max_steps: # 单局步数上限
|
||||
pos, _ = env.reset()
|
||||
total_r, terminated, truncated = 0.0, False, False
|
||||
while not (terminated or truncated): # terminated 到终点 / truncated 超时
|
||||
|
||||
# ε-greedy选动作 explore or exploit
|
||||
if rng.random() < epsilon:
|
||||
@@ -145,17 +34,15 @@ def train(env, episodes=500, alpha=0.1,
|
||||
else:
|
||||
a = int(Q[pos].argmax())
|
||||
|
||||
nxt, r, done = env.step(a)
|
||||
nxt, r, terminated, truncated, _ = env.step(a)
|
||||
# Q-Learning 更新;终点是终止状态,目标值仅为 r
|
||||
Q[pos][a] += alpha * (r if done else r + gamma * Q[nxt].max() - Q[pos][a])
|
||||
Q[pos][a] += alpha * (r if terminated else r + gamma * Q[nxt].max() - Q[pos][a])
|
||||
pos = nxt
|
||||
total_r += r
|
||||
if done:
|
||||
break
|
||||
|
||||
epsilon = max(eps_min, epsilon * eps_decay) # 逐渐减少探索
|
||||
rewards.append(total_r)
|
||||
successes.append(done)
|
||||
successes.append(terminated)
|
||||
ep_steps.append(env.steps)
|
||||
eps_hist.append(epsilon)
|
||||
if si < len(snap_at) and ep == snap_at[si]:
|
||||
@@ -165,156 +52,6 @@ def train(env, episodes=500, alpha=0.1,
|
||||
return Q, rewards, successes, ep_steps, eps_hist, snapshots
|
||||
|
||||
|
||||
def evaluate(env, Q, episodes=100):
|
||||
"""纯贪心策略(无探索)评测, 返回 (成功率, 成功局平均步数)。"""
|
||||
wins, steps_list = 0, []
|
||||
for _ in range(episodes):
|
||||
pos, done = env.reset(), False
|
||||
while not done and env.steps < env.max_steps:
|
||||
pos, _, done = env.step(int(Q[pos].argmax()))
|
||||
if done:
|
||||
wins += 1
|
||||
steps_list.append(env.steps)
|
||||
avg_steps = float(np.mean(steps_list)) if steps_list else float("nan")
|
||||
return wins / episodes, avg_steps
|
||||
|
||||
|
||||
def show_maze(env):
|
||||
print(f"迷宫 {env.shape[0]}x{env.shape[1]}(S起点 G终点 #墙):")
|
||||
for r in range(env.shape[0]):
|
||||
print("".join(CHARS[env.maze[r][c]] for c in range(env.shape[1])))
|
||||
|
||||
|
||||
def show_policy(env, Q):
|
||||
print("学到的策略(每格最优动作):")
|
||||
for r in range(env.shape[0]):
|
||||
row = ""
|
||||
for c in range(env.shape[1]):
|
||||
if env.maze[r][c] == WALL:
|
||||
row += " # "
|
||||
elif (r, c) == env.goal:
|
||||
row += " G "
|
||||
else:
|
||||
row += f" {ARROWS[int(Q[r, c].argmax())]} "
|
||||
print(row)
|
||||
|
||||
|
||||
def show_path(env, Q):
|
||||
print("\n贪心策略走出的一条路径:")
|
||||
pos, path, done = env.reset(), [env.start_pos], False
|
||||
while not done and env.steps < env.max_steps:
|
||||
pos, _, done = env.step(int(Q[pos].argmax()))
|
||||
path.append(pos)
|
||||
grid = [[CHARS[env.maze[r][c]] for c in range(env.shape[1])]
|
||||
for r in range(env.shape[0])]
|
||||
for i, (r, c) in enumerate(path):
|
||||
if (r, c) not in (env.start_pos, env.goal):
|
||||
grid[r][c] = str(i % 10) # 若重复经过同一格,显示最后一步序号
|
||||
print("\n".join(" ".join(row) for row in grid))
|
||||
print(f"到达终点: {'是' if done else '否'},共 {len(path) - 1} 步")
|
||||
|
||||
|
||||
# ─────────────────── 训练过程可视化(需 matplotlib) ───────────────────
|
||||
def _setup_plt():
|
||||
import matplotlib
|
||||
|
||||
matplotlib.use("Agg") # 不弹窗,直接保存图片
|
||||
import matplotlib.pyplot as plt
|
||||
plt.rcParams["font.sans-serif"] = ["Microsoft YaHei", "SimHei"] # 中文字体
|
||||
plt.rcParams["axes.unicode_minus"] = False # 正常显示负号
|
||||
return plt
|
||||
|
||||
|
||||
def moving_avg(x, k=20):
|
||||
x = np.asarray(x, dtype=float)
|
||||
if len(x) < k:
|
||||
return np.array([])
|
||||
return np.convolve(x, np.ones(k) / k, mode="valid")
|
||||
|
||||
|
||||
def _plot_ma(ax, x, window, color, label):
|
||||
x = np.asarray(x, dtype=float)
|
||||
ax.plot(x, color="gray", alpha=0.35, label="逐局值")
|
||||
ma = moving_avg(x, window)
|
||||
if len(ma):
|
||||
ax.plot(np.arange(window - 1, window - 1 + len(ma)), ma,
|
||||
color=color, lw=2, label=f"滑动平均(窗口{window})")
|
||||
ax.set_xlabel("训练局数 Episode")
|
||||
ax.legend()
|
||||
|
||||
|
||||
def plot_training(env, rewards, successes, ep_steps, eps_hist,
|
||||
window=20, out="qlearning_training.png"):
|
||||
"""四联图: 回报曲线 / 滑动成功率 / 单局步数 / ε 衰减。"""
|
||||
plt = _setup_plt()
|
||||
fig, axes = plt.subplots(2, 2, figsize=(12, 8))
|
||||
fig.suptitle("Q-Learning 走迷宫训练过程", fontsize=15)
|
||||
|
||||
ax = axes[0][0]
|
||||
_plot_ma(ax, rewards, window, "C0", None)
|
||||
ax.set_ylabel("单局回报")
|
||||
ax.set_title("逐局回报(越高越好)")
|
||||
|
||||
ax = axes[0][1]
|
||||
_plot_ma(ax, successes, window, "C1", None)
|
||||
ax.set_ylim(-0.05, 1.05)
|
||||
ax.set_ylabel("成功率")
|
||||
ax.set_title(f"滑动成功率(窗口{window})")
|
||||
|
||||
ax = axes[1][0]
|
||||
_plot_ma(ax, ep_steps, window, "C2", None)
|
||||
sp = shortest_path_len(env)
|
||||
if sp >= 0:
|
||||
ax.axhline(sp, color="red", ls="--", lw=1.5, label=f"最短路 {sp} 步")
|
||||
ax.set_ylabel("单局步数")
|
||||
ax.set_title("单局步数(降到红线 = 学会最短路)")
|
||||
ax.legend()
|
||||
|
||||
ax = axes[1][1]
|
||||
ax.plot(eps_hist, color="C3", lw=2)
|
||||
ax.set_xlabel("训练局数 Episode")
|
||||
ax.set_ylabel("ε")
|
||||
ax.set_title(f"探索率 ε 衰减({eps_hist[0]:.2f} → {eps_hist[-1]:.2f})")
|
||||
|
||||
fig.tight_layout(rect=[0, 0, 1, 0.96])
|
||||
fig.savefig(out, dpi=150)
|
||||
print(f"\n图片已保存: {out}")
|
||||
|
||||
|
||||
def plot_snapshots(env, snapshots, Q_final, out="qlearning_snapshots.png"):
|
||||
"""V(s)=max Q 热力图随训练的变化:观察价值从终点 G 逐步回传到起点。"""
|
||||
plt = _setup_plt()
|
||||
panels = snapshots + [("最终", Q_final)]
|
||||
vmin = min(float(q.max(axis=2).min()) for _, q in panels)
|
||||
vmax = max(float(q.max(axis=2).max()) for _, q in panels)
|
||||
if vmin == vmax:
|
||||
vmin, vmax = vmin - 1, vmax + 1
|
||||
|
||||
n = len(panels)
|
||||
fig, axes = plt.subplots(1, n, figsize=(3 * n, 3.4))
|
||||
axes = np.atleast_1d(axes)
|
||||
masked_wall = env.maze == WALL
|
||||
for ax, (tag, q) in zip(axes, panels):
|
||||
V = np.ma.masked_where(masked_wall, q.max(axis=2))
|
||||
im = ax.imshow(V, cmap="viridis", vmin=vmin, vmax=vmax, origin="upper")
|
||||
fs = 8 if max(env.shape) <= 7 else 5
|
||||
for r in range(env.shape[0]):
|
||||
for c in range(env.shape[1]):
|
||||
if masked_wall[r, c]:
|
||||
ax.text(c, r, "█", ha="center", va="center",
|
||||
color="black", fontsize=fs)
|
||||
else:
|
||||
ax.text(c, r, f"{V[r, c]:.0f}", ha="center", va="center",
|
||||
color="white", fontsize=fs)
|
||||
ax.set_xticks([])
|
||||
ax.set_yticks([])
|
||||
ax.set_title(f"第 {tag} 局后" if tag != "最终" else "训练结束")
|
||||
fig.suptitle("状态价值 V(s)=max Q 的学习过程(价值从 G 逐步回传)", fontsize=13)
|
||||
fig.colorbar(im, ax=axes, fraction=0.03, pad=0.02)
|
||||
fig.savefig(out, dpi=150, bbox_inches="tight")
|
||||
print(f"图片已保存: {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Q-Learning 走迷宫(numpy 表格型实现)")
|
||||
parser.add_argument("--seed", type=int, default=0, help="迷宫生成随机种子")
|
||||
@@ -332,13 +69,22 @@ if __name__ == "__main__":
|
||||
snap_at = (1, 10, 50, 200, 500) if args.plot else ()
|
||||
Q, rewards, successes, ep_steps, eps_hist, snapshots = train(
|
||||
env, episodes=args.episodes, snap_at=snap_at)
|
||||
show_policy(env, Q)
|
||||
show_path(env, Q)
|
||||
rate, avg_steps = evaluate(env, Q)
|
||||
|
||||
def act(pos):
|
||||
return int(Q[pos].argmax())
|
||||
|
||||
show_policy(env, act)
|
||||
show_path(env, act)
|
||||
rate, avg_steps = evaluate(env, act)
|
||||
print(f"\n前50局平均回报: {np.mean(rewards[:50]):6.1f} "
|
||||
f"后50局平均回报: {np.mean(rewards[-50:]):6.1f}")
|
||||
print(f"后50局成功率: {np.mean(successes[-50:]):.0%}")
|
||||
print(f"评测(100局贪心): 成功率 {rate:.0%},平均 {avg_steps:.1f} 步到达终点")
|
||||
if args.plot:
|
||||
plot_training(env, rewards, successes, ep_steps, eps_hist)
|
||||
plot_snapshots(env, snapshots, Q)
|
||||
plot_training(env, rewards, successes, ep_steps,
|
||||
extra=(eps_hist, "ε",
|
||||
f"探索率 ε 衰减({eps_hist[0]:.2f} → {eps_hist[-1]:.2f})", False),
|
||||
out="qlearning_training.png", title="Q-Learning 走迷宫训练过程")
|
||||
plot_snapshots(env, [(ep, q.max(axis=2)) for ep, q in snapshots], Q.max(axis=2),
|
||||
out="qlearning_snapshots.png",
|
||||
title="状态价值 V(s)=max Q 的学习过程(价值从 G 逐步回传)")
|
||||
|
||||
Reference in New Issue
Block a user