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