"""CartPole-v1：随机策略基线，建立环境直觉。"""
import numpy as np
import gymnasium as gym
import matplotlib.pyplot as plt


def run_random_agent(n_episodes: int = 200, seed: int = 42):
    env = gym.make("CartPole-v1")
    env.reset(seed=seed)
    returns = []

    for ep in range(n_episodes):
        obs, info = env.reset()
        done = False
        truncated = False
        ep_return = 0.0

        while not (done or truncated):
            # 完全随机探索 —— 无任何学习
            action = env.action_space.sample()
            obs, reward, done, truncated, info = env.step(action)
            ep_return += reward

        returns.append(ep_return)
        if (ep + 1) % 20 == 0:
            print(
                f"Episode {ep+1:3d} | Return = {ep_return:6.1f} | "
                f"Mean(last20) = {np.mean(returns[-20:]):6.1f}"
            )

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


if __name__ == "__main__":
    rets = run_random_agent(n_episodes=200)

    window = 20
    smooth = np.convolve(rets, np.ones(window) / window, mode="valid")

    plt.figure(figsize=(8, 4))
    plt.plot(rets, alpha=0.35, 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 · Random Agent Baseline")
    plt.legend()
    plt.grid(True, alpha=0.3)
    plt.tight_layout()
    plt.savefig("cartpole_random_baseline.png", dpi=150)
    print("曲线已保存: cartpole_random_baseline.png")
    print(f"全程平均回报: {rets.mean():.1f} ± {rets.std():.1f}")
