Views
No views yet
1import gymnasium as gym
2from time import sleep
3from huggingface_sb3 import package_to_hub
4from stable_baselines3 import PPO
5from stable_baselines3.common.env_util import make_vec_env
6from stable_baselines3.common.evaluation import evaluate_policy
7from stable_baselines3.common.monitor import Monitor
8from stable_baselines3.common.vec_env import DummyVecEnv
9
10# Create the environment
11env = make_vec_env("LunarLander-v2", n_envs=16)
12
13# We added some parameters to accelerate the training
14model = PPO(
15 policy="MlpPolicy",
16 env=env,
17 n_steps=1024,
18 batch_size=64,
19 n_epochs=4,
20 gamma=0.999,
21 gae_lambda=0.98,
22 ent_coef=0.01,
23 verbose=1,
24)
25
26# Train it for 1,000,000 timesteps
27model.learn(total_timesteps=1000000)
28# Save the model
29model.save(model_name)
30
31# Test the model
32# model = PPO.load(model_name)
33eval_env = Monitor(gym.make("LunarLander-v2"))
34mean_reward, std_reward = evaluate_policy(model, eval_env, n_eval_episodes=10, deterministic=True)
35print(f"mean_reward={mean_reward:.2f} +/- {std_reward}")
36
37# Visualize the model
38env = gym.make("LunarLander-v2", render_mode='human')
39
40state, _ = env.reset()
41stop = False
42
43while not stop:
44 action, _ = model.predict(state)
45 state, reward, terminated, truncated, info = env.step(action)
46 stop = terminated or truncated
47 env.render()
48 sleep(0.05)
49
50 if terminated or truncated:
51 observation, info = env.reset()
52
53env.close()
54...