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

レッスン 9: ジムと安定したベースライン 3 — 現実の RL フレームワーク

体育館 API の詳細。ラッパー。 Stable-Baselines3 のトレーニング、評価、コールバック。ハイパーパラメータの調整。 TensorBoard のロギング。

🧠 AI と ML — レッスン 8 レッスン 9: 体育館と安定したベースライン 3 — Rs 実戦フレームワーク

強化学習: 基礎から高度まで

パート 3: RL フレームワークと実践

xdev.asia

はじめに

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ベイジアン ハイパーパラメータ調整