はじめに
Gymnasium (後継 OpenAI Gym) は、RL 環境の標準 API です。 Stable-Baselines3 (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 --logdir ./tensorboard/
# Monitor: reward, loss, learning rate, etc.
概要
| フレームワーク | 役割 | 主な機能 |
|---|---|---|
| 体育館 | 環境API | 標準化、ラッパー |
| SB3 | アルゴリズムライブラリ | PPO、SAC、DQN — 実稼働対応 |
| テンソルボード | ビジュアライゼーション | リアルタイムのトレーニング指標 |
| オプチュナ | HPO | ベイジアン ハイパーパラメータ調整 |