"""教学向：状态离散化 + ε-greedy Q-learning 骨架（非生产级）。"""
import numpy as np
import gymnasium as gym
import matplotlib.pyplot as plt


def discretize(obs, bins):
    """把连续 obs 映射到离散索引元组。"""
    indices = []
    # obs: [x, x_dot, theta, theta_dot]
    ranges = [
        (-2.4, 2.4),
        (-3.0, 3.0),
        (-0.2, 0.2),
        (-3.0, 3.0),
    ]
    for o, (lo, hi), b in zip(obs, ranges, bins):
        o_clip = np.clip(o, lo, hi)
        idx = int((o_clip - lo) / (hi - lo) * (b - 1))
        indices.append(idx)
    return tuple(indices)


def q_learning_cartpole(
    n_episodes=2000,
    bins=(6, 6, 12, 12),
    alpha=0.1,
    gamma=0.99,
    eps_start=1.0,
    eps_end=0.05,
    eps_decay_episodes=1500,
):
    env = gym.make("CartPole-v1")
    n_actions = env.action_space.n
    Q = np.zeros(bins + (n_actions,))
    returns = []

    for ep in range(n_episodes):
        frac = min(1.0, ep / eps_decay_episodes)
        eps = eps_start + frac * (eps_end - eps_start)

        obs, _ = env.reset()
        s = discretize(obs, bins)
        done = truncated = False
        ep_ret = 0.0

        while not (done or truncated):
            if np.random.rand() < eps:
                a = env.action_space.sample()
            else:
                a = int(np.argmax(Q[s]))

            next_obs, r, done, truncated, _ = env.step(a)
            s_next = discretize(next_obs, bins)
            ep_ret += r

            # Q-learning 更新
            td_target = r + gamma * np.max(Q[s_next]) * (not (done or truncated))
            Q[s + (a,)] += alpha * (td_target - Q[s + (a,)])
            s = s_next

        returns.append(ep_ret)
        if (ep + 1) % 100 == 0:
            print(
                f"Ep {ep+1:4d} | eps={eps:.3f} | "
                f"mean100={np.mean(returns[-100:]):6.1f}"
            )

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


if __name__ == "__main__":
    rets, Q = q_learning_cartpole()
    print("最终 100 局平均:", rets[-100:].mean())

    window = 50
    smooth = np.convolve(rets, np.ones(window) / window, mode="valid")
    plt.figure(figsize=(8, 4))
    plt.plot(rets, alpha=0.25, 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 · Discretized Q-learning")
    plt.legend()
    plt.grid(True, alpha=0.3)
    plt.tight_layout()
    plt.savefig("cartpole_q_learning.png", dpi=150)
    print("曲线已保存: cartpole_q_learning.png")
