"""PPO · CartPole-v1 一键版（Stable-Baselines3）。

适合先「看见收敛」，再回头抠 PPO clip / GAE 公式。
依赖: pip install stable-baselines3 gymnasium matplotlib
"""
from __future__ import annotations

import gymnasium as gym
import matplotlib.pyplot as plt
import numpy as np
from stable_baselines3 import PPO
from stable_baselines3.common.callbacks import BaseCallback
from stable_baselines3.common.monitor import Monitor


class ReturnTracker(BaseCallback):
    def __init__(self):
        super().__init__()
        self.returns = []

    def _on_step(self) -> bool:
        for info in self.locals.get("infos", []):
            if "episode" in info:
                self.returns.append(info["episode"]["r"])
        return True


def train(total_timesteps: int = 50_000, seed: int = 42):
    env = Monitor(gym.make("CartPole-v1"))
    model = PPO(
        "MlpPolicy",
        env,
        verbose=0,
        seed=seed,
        learning_rate=3e-4,
        n_steps=2048,
        batch_size=64,
        gamma=0.99,
        gae_lambda=0.95,
        clip_range=0.2,
        ent_coef=0.0,
    )
    tracker = ReturnTracker()
    model.learn(total_timesteps=total_timesteps, callback=tracker)
    env.close()
    return model, np.array(tracker.returns)


if __name__ == "__main__":
    model, rets = train()
    if len(rets) == 0:
        print("未记录到 episode 回报，请检查 Monitor 包装。")
    else:
        print(f"共 {len(rets)} 个 episode")
        print(f"最终 20 局平均回报: {rets[-20:].mean():.1f}")

        window = 20
        if len(rets) >= window:
            smooth = np.convolve(rets, np.ones(window) / window, mode="valid")
            xs = range(window - 1, len(rets))
        else:
            smooth, xs = rets, range(len(rets))

        plt.figure(figsize=(8, 4))
        plt.plot(rets, alpha=0.35, label="Episode return")
        plt.plot(xs, smooth, label=f"Moving avg (w={window})")
        plt.xlabel("Episode")
        plt.ylabel("Return")
        plt.title("CartPole-v1 · PPO (Stable-Baselines3)")
        plt.legend()
        plt.grid(True, alpha=0.3)
        plt.tight_layout()
        plt.savefig("ppo_cartpole_returns.png", dpi=150)
        print("曲线已保存: ppo_cartpole_returns.png")

    model.save("ppo_cartpole_sb3")
    print("模型已保存: ppo_cartpole_sb3.zip")
