"""最小可运行 DQN · CartPole-v1（PyTorch）。

核心技巧:
  1) Replay Buffer  —— 打乱样本相关性
  2) Target Network —— 稳定 TD 目标
依赖: pip install torch gymnasium matplotlib numpy
"""
from __future__ import annotations

import random
from collections import deque
from dataclasses import dataclass

import gymnasium as gym
import matplotlib.pyplot as plt
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F


@dataclass
class Config:
    gamma: float = 0.99
    lr: float = 1e-3
    batch_size: int = 64
    buffer_size: int = 50_000
    min_buffer: int = 1_000
    target_sync: int = 200
    eps_start: float = 1.0
    eps_end: float = 0.05
    eps_decay_steps: int = 10_000
    total_steps: int = 40_000
    hidden: int = 128
    seed: int = 42
    device: str = "cpu"


class QNet(nn.Module):
    def __init__(self, n_obs: int, n_act: int, hidden: int):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(n_obs, hidden),
            nn.ReLU(),
            nn.Linear(hidden, hidden),
            nn.ReLU(),
            nn.Linear(hidden, n_act),
        )

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.net(x)


class ReplayBuffer:
    def __init__(self, capacity: int):
        self.buf = deque(maxlen=capacity)

    def push(self, s, a, r, s2, done):
        self.buf.append((s, a, r, s2, done))

    def sample(self, batch_size: int):
        batch = random.sample(self.buf, batch_size)
        s, a, r, s2, d = map(np.array, zip(*batch))
        return s, a, r, s2, d

    def __len__(self):
        return len(self.buf)


def epsilon_by_step(step: int, cfg: Config) -> float:
    frac = min(1.0, step / cfg.eps_decay_steps)
    return cfg.eps_start + frac * (cfg.eps_end - cfg.eps_start)


def train(cfg: Config | None = None):
    cfg = cfg or Config()
    random.seed(cfg.seed)
    np.random.seed(cfg.seed)
    torch.manual_seed(cfg.seed)

    env = gym.make("CartPole-v1")
    n_obs = env.observation_space.shape[0]
    n_act = env.action_space.n

    q = QNet(n_obs, n_act, cfg.hidden).to(cfg.device)
    q_tgt = QNet(n_obs, n_act, cfg.hidden).to(cfg.device)
    q_tgt.load_state_dict(q.state_dict())
    opt = torch.optim.Adam(q.parameters(), lr=cfg.lr)
    buf = ReplayBuffer(cfg.buffer_size)

    s, _ = env.reset(seed=cfg.seed)
    ep_ret, returns, losses = 0.0, [], []
    ep_count = 0

    for step in range(1, cfg.total_steps + 1):
        eps = epsilon_by_step(step, cfg)
        if random.random() < eps:
            a = env.action_space.sample()
        else:
            with torch.no_grad():
                qs = q(torch.as_tensor(s, dtype=torch.float32, device=cfg.device).unsqueeze(0))
                a = int(qs.argmax(dim=1).item())

        s2, r, term, trunc, _ = env.step(a)
        done = term or trunc
        buf.push(s, a, r, s2, float(done))
        ep_ret += r
        s = s2

        if done:
            returns.append(ep_ret)
            ep_count += 1
            if ep_count % 20 == 0:
                mean20 = np.mean(returns[-20:])
                print(
                    f"step={step:6d} | ep={ep_count:4d} | "
                    f"ret={ep_ret:6.1f} | mean20={mean20:6.1f} | eps={eps:.3f}"
                )
            s, _ = env.reset()
            ep_ret = 0.0

        if len(buf) < cfg.min_buffer:
            continue

        bs, ba, br, bs2, bd = buf.sample(cfg.batch_size)
        bs_t = torch.as_tensor(bs, dtype=torch.float32, device=cfg.device)
        ba_t = torch.as_tensor(ba, dtype=torch.int64, device=cfg.device)
        br_t = torch.as_tensor(br, dtype=torch.float32, device=cfg.device)
        bs2_t = torch.as_tensor(bs2, dtype=torch.float32, device=cfg.device)
        bd_t = torch.as_tensor(bd, dtype=torch.float32, device=cfg.device)

        q_sa = q(bs_t).gather(1, ba_t.unsqueeze(1)).squeeze(1)
        with torch.no_grad():
            max_next = q_tgt(bs2_t).max(dim=1).values
            y = br_t + cfg.gamma * max_next * (1.0 - bd_t)

        loss = F.mse_loss(q_sa, y)
        opt.zero_grad()
        loss.backward()
        nn.utils.clip_grad_norm_(q.parameters(), 10.0)
        opt.step()
        losses.append(float(loss.item()))

        if step % cfg.target_sync == 0:
            q_tgt.load_state_dict(q.state_dict())

    env.close()
    return np.array(returns), np.array(losses)


if __name__ == "__main__":
    rets, losses = train()
    print(f"最终 50 局平均回报: {rets[-50:].mean():.1f}")

    window = 20
    smooth = np.convolve(rets, np.ones(window) / window, mode="valid")
    plt.figure(figsize=(8, 4))
    plt.plot(rets, alpha=0.3, label="Episode return")
    plt.plot(range(window - 1, len(rets)), smooth, label=f"Moving avg (w={window})")
    plt.xlabel("Episode")
    plt.ylabel("Return")
    plt.title("CartPole-v1 · Minimal DQN")
    plt.legend()
    plt.grid(True, alpha=0.3)
    plt.tight_layout()
    plt.savefig("dqn_cartpole_returns.png", dpi=150)
    print("曲线已保存: dqn_cartpole_returns.png")
