91 lines
3.8 KiB
Python
91 lines
3.8 KiB
Python
import argparse
|
||
|
||
import numpy as np
|
||
|
||
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,
|
||
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, _ = 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:
|
||
a = int(rng.integers(env.n_actions))
|
||
else:
|
||
a = int(Q[pos].argmax())
|
||
|
||
nxt, r, terminated, truncated, _ = env.step(a)
|
||
# Q-Learning 更新;终点是终止状态,目标值仅为 r
|
||
Q[pos][a] += alpha * (r if terminated else r + gamma * Q[nxt].max() - Q[pos][a])
|
||
pos = nxt
|
||
total_r += r
|
||
|
||
epsilon = max(eps_min, epsilon * eps_decay) # 逐渐减少探索
|
||
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, Q.copy()))
|
||
si += 1
|
||
|
||
return Q, rewards, successes, ep_steps, eps_hist, snapshots
|
||
|
||
|
||
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)
|
||
|
||
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,
|
||
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 逐步回传)")
|