first commit
This commit is contained in:
@@ -0,0 +1,344 @@
|
||||
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
|
||||
|
||||
|
||||
def train(env, episodes=500, alpha=0.1,
|
||||
gamma=0.9, epsilon=1.0, eps_min=0.05, eps_decay=0.99, seed=0, snap_at=()):
|
||||
rng = np.random.default_rng(seed)
|
||||
Q = np.zeros((*env.shape, env.n_actions)) # Q_{s, a} 这里是 Q[r, c, a]
|
||||
rewards, successes, ep_steps, eps_hist = [], [], [], []
|
||||
snapshots = [] # [(局数, Q副本), ...]
|
||||
snap_at = sorted({int(e) for e in snap_at if 0 < int(e) <= episodes})
|
||||
si = 0
|
||||
|
||||
for ep in range(1, episodes + 1):
|
||||
pos, total_r, done = env.reset(), 0.0, False
|
||||
while env.steps < env.max_steps: # 单局步数上限
|
||||
|
||||
# ε-greedy选动作 explore or exploit
|
||||
if rng.random() < epsilon:
|
||||
a = int(rng.integers(env.n_actions))
|
||||
else:
|
||||
a = int(Q[pos].argmax())
|
||||
|
||||
nxt, r, done = env.step(a)
|
||||
# Q-Learning 更新;终点是终止状态,目标值仅为 r
|
||||
Q[pos][a] += alpha * (r if done 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)
|
||||
ep_steps.append(env.steps)
|
||||
eps_hist.append(epsilon)
|
||||
if si < len(snap_at) and ep == snap_at[si]:
|
||||
snapshots.append((ep, Q.copy()))
|
||||
si += 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="迷宫生成随机种子")
|
||||
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=500, help="训练局数")
|
||||
parser.add_argument("--braid", type=float, default=0.0,
|
||||
help="拆墙成环概率 0~1, 越大岔路/环路越多")
|
||||
parser.add_argument("--plot", action="store_true",
|
||||
help="生成 qlearning_training.png 与 qlearning_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, 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)
|
||||
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)
|
||||
Reference in New Issue
Block a user