# Source: content/notes/deep-learning/deep-rl.md
# Independent CPU example; use the curriculum environment.
# See /notes/ml/#example-environment or /notes/deep-learning/#example-environment.

import numpy as np
import torch
import gymnasium as gym
from stable_baselines3 import PPO
from stable_baselines3.common.env_util import make_vec_env

torch.set_num_threads(1)
training_seed = 63
evaluation_seeds = list(range(9000, 9008))
train_env = make_vec_env("CartPole-v1", n_envs=4, seed=training_seed)
try:
    model = PPO("MlpPolicy", train_env, device="cpu", seed=training_seed,
                n_steps=128, batch_size=128, n_epochs=6,
                learning_rate=3e-4, gamma=0.99, gae_lambda=0.95,
                clip_range=0.2, verbose=0)
    initial_action_weights = model.policy.action_net.weight.detach().clone()
    model.learn(total_timesteps=12800, progress_bar=False)
    assert model.num_timesteps == 12800
    assert not torch.equal(model.policy.action_net.weight, initial_action_weights)
    assert all(torch.isfinite(p).all() for p in model.policy.parameters())
finally:
    train_env.close()

def evaluate(policy):
    env = gym.make("CartPole-v1")
    returns, endings = [], []
    try:
        for seed in evaluation_seeds:
            observation, info = env.reset(seed=seed)
            env.action_space.seed(seed + 10000)
            total_reward = 0.0
            for step in range(env.spec.max_episode_steps):
                if policy is None:
                    action = env.action_space.sample()
                else:
                    action, _ = policy.predict(observation, deterministic=True)
                    action = int(np.asarray(action).item())
                observation, reward, terminated, truncated, info = env.step(action)
                assert np.isfinite(observation).all() and np.isfinite(reward)
                total_reward += float(reward)
                if terminated or truncated:
                    endings.append((bool(terminated), bool(truncated)))
                    break
            else:
                raise AssertionError("The environment did not signal its time limit")
            returns.append(total_reward)
    finally:
        env.close()
    returns = np.asarray(returns)
    assert len(returns) == len(evaluation_seeds)
    assert np.isfinite(returns).all() and np.all((returns >= 1) & (returns <= 500))
    return {"episode_returns": returns.tolist(), "mean": float(returns.mean()),
            "sample_std": float(returns.std(ddof=1)),
            "terminated": sum(end[0] for end in endings),
            "truncated": sum(end[1] for end in endings)}

random_result = evaluate(None)
ppo_result = evaluate(model)
print({"environment_steps": model.num_timesteps,
       "evaluation_seeds": evaluation_seeds,
       "random_policy": random_result, "deterministic_ppo": ppo_result})
