Files
RL-learning/q_learning/q_learning.py
T

91 lines
3.8 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
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 逐步回传)")