はじめに
DQN (Deep Q-Network) — DeepMind が Q テーブルの代わりにニューラル ネットワークを使用し、Atari ゲームで超人的なパフォーマンスを達成したときの大きな飛躍 (Nature、2015)。
1. Q テーブルから DQN へ
| Qテーブル | DQN |
|---|---|
| Q(s,a) をテーブルに格納 | ニューラル ネットワークは Q(s,a;θ) |
| 離散的な小さな状態空間でのみ機能します。連続した高次元の状態を扱う | |
| 正確 | おおよそ |
2. DQN アーキテクチャ
import torch
import torch.nn as nn
class DQN(nn.Module):
def __init__(self, state_dim, action_dim, hidden_dim=128):
super().__init__()
self.net = nn.Sequential(
nn.Linear(state_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, action_dim),
)
def forward(self, x):
return self.net(x)
3. 体験リプレイ
from collections import deque
import random
class ReplayBuffer:
def __init__(self, capacity=100_000):
self.buffer = deque(maxlen=capacity)
def push(self, state, action, reward, next_state, done):
self.buffer.append((state, action, reward, next_state, done))
def sample(self, batch_size):
batch = random.sample(self.buffer, batch_size)
states, actions, rewards, next_states, dones = zip(*batch)
return (
torch.FloatTensor(np.array(states)),
torch.LongTensor(actions),
torch.FloatTensor(rewards),
torch.FloatTensor(np.array(next_states)),
torch.FloatTensor(dones),
)
def __len__(self):
return len(self.buffer)
4. DQN トレーニング ループ
class DQNAgent:
def __init__(self, state_dim, action_dim, lr=1e-3, gamma=0.99):
self.online_net = DQN(state_dim, action_dim)
self.target_net = DQN(state_dim, action_dim)
self.target_net.load_state_dict(self.online_net.state_dict())
self.optimizer = torch.optim.Adam(self.online_net.parameters(), lr=lr)
self.buffer = ReplayBuffer()
self.gamma = gamma
def update(self, batch_size=64):
states, actions, rewards, next_states, dones = self.buffer.sample(batch_size)
# Current Q values
q_values = self.online_net(states).gather(1, actions.unsqueeze(1)).squeeze()
# Target Q values
with torch.no_grad():
next_q = self.target_net(next_states).max(1)[0]
target = rewards + self.gamma * next_q * (1 - dones)
loss = nn.MSELoss()(q_values, target)
self.optimizer.zero_grad()
loss.backward()
self.optimizer.step()
def sync_target(self):
self.target_net.load_state_dict(self.online_net.state_dict())
5. DQN の改善
ダブル DQN
アクションの選択と評価を分離 → 過大評価を減らす:
# Double DQN target
best_actions = self.online_net(next_states).argmax(1)
next_q = self.target_net(next_states).gather(1, best_actions.unsqueeze(1)).squeeze()
決闘DQN
値 V(s) と利点 A(s,a) を分離します。
class DuelingDQN(nn.Module):
def __init__(self, state_dim, action_dim):
super().__init__()
self.feature = nn.Sequential(nn.Linear(state_dim, 128), nn.ReLU())
self.value = nn.Sequential(nn.Linear(128, 64), nn.ReLU(), nn.Linear(64, 1))
self.advantage = nn.Sequential(nn.Linear(128, 64), nn.ReLU(), nn.Linear(64, action_dim))
def forward(self, x):
feat = self.feature(x)
val = self.value(feat)
adv = self.advantage(feat)
return val + adv - adv.mean(dim=1, keepdim=True)
概要
| 改善点 | 問題解決 |
|---|---|
| 体験リプレイ | 連続サンプル間の相関 |
| ターゲットネットワーク | トレーニングの不安定性 (ターゲットの移動) |
| ダブルDQN | Q値の過大評価 |
| DQNとの決闘 | より良い価値推定 |
| 優先再生 | 重要な経験のサンプルをもっと見る |
| レインボー | すべての改善点を組み合わせる |