Files
RL-learning/dqn/dqn_maze.py
T

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