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 回传)")