Files
RL-learning/q_learning/q_learning.py
T
2026-08-31 17:55:15 +08:00

345 lines
13 KiB
Python
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.
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)