強化学習: 基礎から高度まで
パート 3: RL フレームワークと実践
xdev.asia
はじめに
状態/アクション空間の設計から報酬の形成やエージェントのトレーニングに至るまで、独自の問題に合わせてカスタムのジム環境を構築します。
1. カスタム環境テンプレート
import gymnasium as gym
from gymnasium import spaces
import numpy as np
class SnakeEnv(gym.Env):
metadata = {"render_modes": ["human", "rgb_array"], "render_fps": 10}
def __init__(self, grid_size=10, render_mode=None):
super().__init__()
self.grid_size = grid_size
self.render_mode = render_mode
# Action: 0=up, 1=right, 2=down, 3=left
self.action_space = spaces.Discrete(4)
# Observation: grid with snake body, head, food
self.observation_space = spaces.Box(
low=0, high=3, shape=(grid_size, grid_size), dtype=np.uint8
)
def reset(self, seed=None, options=None):
super().reset(seed=seed)
self.snake = [(5, 5)]
self.direction = 1 # right
self.food = self._place_food()
self.score = 0
self.steps = 0
return self._get_obs(), self._get_info()
def step(self, action):
self.steps += 1
self._move_snake(action)
terminated = self._check_collision()
truncated = self.steps >= self.grid_size * self.grid_size * 2
reward = self._compute_reward(terminated)
return self._get_obs(), reward, terminated, truncated, self._get_info()
def _place_food(self):
while True:
pos = (self.np_random.integers(0, self.grid_size),
self.np_random.integers(0, self.grid_size))
if pos not in self.snake:
return pos
def _get_obs(self):
grid = np.zeros((self.grid_size, self.grid_size), dtype=np.uint8)
for segment in self.snake:
grid[segment] = 1 # body
grid[self.snake[0]] = 2 # head
grid[self.food] = 3 # food
return grid
def _get_info(self):
return {"score": self.score, "length": len(self.snake)}
2. 報酬の形成
def _compute_reward(self, terminated):
if terminated:
return -10.0
if self.snake[0] == self.food:
self.score += 1
return 10.0
# Distance-based shaping
head = self.snake[0]
dist_to_food = abs(head[0] - self.food[0]) + abs(head[1] - self.food[1])
prev_dist = abs(self.prev_head[0] - self.food[0]) + abs(self.prev_head[1] - self.food[1])
if dist_to_food < prev_dist:
return 0.1 # Moving closer
else:
return -0.1 # Moving away
報酬設計の原則
| 原則 | 説明 |
|---|
| 疎 vs 密 | 高密度の報酬は学習を早めますが、報酬のハッキングは危険です |
| 大きさ | プラス/マイナスの報酬のバランスをとる |
| 整形 | 最適なポリシーを変更せずにエージェントをガイド |
| ポテンシャルベース | ポリシーの不変性を保証 |
3. 登録とトレーニング
# Register custom env
gym.register(id="Snake-v0", entry_point="snake_env:SnakeEnv")
# Train with SB3
from stable_baselines3 import PPO
env = gym.make("Snake-v0")
model = PPO("MlpPolicy", env, verbose=1)
model.learn(total_timesteps=500_000)
概要
| 側面 | ベストプラクティス |
|---|
| 観察 | 最小限、有益、正規化された |
| アクションスペース | 可能な場合はディスクリート |
| 報酬 | 密な造形 + 疎なボーナス |
| 終了 | 晴天、公正な条件 |
| テスト | 最初にランダムなエージェントで確認します |