Chuyển đến nội dung chính

第 9 課:Gymnasium 與穩定基線 3 — 現實生活中的 RL 框架

體育館 API 詳細資料。包裝紙。穩定基線3 訓練、評估、回調。超參數調整。 TensorBoard 日誌記錄。

🧠 人工智慧與機器學習 — 第 8 課 第 9 課:體育館與穩定基線3 — 盧比 實戰框架

強化學習:從基礎到高級

第 3 部分:強化學習架構與實踐

亞洲開發網

簡介

Gymnasium(OpenAI Gym 的後繼者)是 RL 環境的標準 API。 穩定基線3 (SB3) 提供 DQN、PPO、SAC、TD3 的生產就緒實施。


1. 體操基礎知識

import gymnasium as gym

# Create environment
env = gym.make("CartPole-v1", render_mode="human")
obs, info = env.reset(seed=42)

for step in range(1000):
    action = env.action_space.sample()  # Random policy
    obs, reward, terminated, truncated, info = env.step(action)
    
    if terminated or truncated:
        obs, info = env.reset()

env.close()

空格

# Discrete actions
env = gym.make("CartPole-v1")
print(env.action_space)  # Discrete(2)

# Continuous actions
env = gym.make("Pendulum-v1")
print(env.action_space)  # Box(-2.0, 2.0, (1,))

# Complex observations
env = gym.make("CarRacing-v2")
print(env.observation_space)  # Box(0, 255, (96, 96, 3))

2. 包裝器

from gymnasium.wrappers import TimeLimit, RecordVideo, NormalizeObservation

env = gym.make("CartPole-v1")
env = TimeLimit(env, max_episode_steps=500)
env = NormalizeObservation(env)
env = RecordVideo(env, "videos/", episode_trigger=lambda e: e % 100 == 0)

3. 穩定基線3 訓練

from stable_baselines3 import PPO, DQN, SAC
from stable_baselines3.common.env_util import make_vec_env

# Vectorized environments for faster training
env = make_vec_env("CartPole-v1", n_envs=4)

model = PPO("MlpPolicy", env, verbose=1,
    learning_rate=3e-4,
    n_steps=2048,
    batch_size=64,
    n_epochs=10,
    gamma=0.99,
    tensorboard_log="./tensorboard/"
)
model.learn(total_timesteps=500_000)
model.save("ppo_cartpole")

4. 評估

from stable_baselines3.common.evaluation import evaluate_policy

model = PPO.load("ppo_cartpole")
eval_env = gym.make("CartPole-v1")
mean_reward, std = evaluate_policy(model, eval_env, n_eval_episodes=20)
print(f"Mean reward: {mean_reward:.2f} +/- {std:.2f}")

5. 回調

from stable_baselines3.common.callbacks import EvalCallback, CheckpointCallback

eval_callback = EvalCallback(
    eval_env, best_model_save_path="./best/",
    log_path="./logs/", eval_freq=5000,
    n_eval_episodes=10, deterministic=True,
)

checkpoint_callback = CheckpointCallback(
    save_freq=10000, save_path="./checkpoints/"
)

model.learn(
    total_timesteps=500_000,
    callback=[eval_callback, checkpoint_callback]
)

6. TensorBoard

tensorboard --logdir ./tensorboard/
# Monitor: reward, loss, learning rate, etc.

總結

框架角色主要特點
體育館環境API標準化,包裝紙
SB3演算法庫PPO、SAC、DQN — 生產就緒
張量板視覺化即時訓練指標
奧圖納磷酸二氫鉀貝葉斯超參數調整