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

レッスン 5: DQN — ディープ Q ネットワークとエクスペリエンス リプレイ

ニューラルネットワークによる関数近似。 DQN アーキテクチャ。リプレイバッファーを体験してください。ターゲットネットワーク。ダブルDQN、決闘DQN、レインボーDQN。

🧠 AI と ML — レッスン 4 レッスン 5: DQN — ディープ Q ネットワークと経験 リプレイ

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

パート 2: 深層強化学習 — ニューラル ネットワークと RL の出会い

xdev.asia

はじめに

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)

概要

改善点問題解決
体験リプレイ連続サンプル間の相関
ターゲットネットワークトレーニングの不安定性 (ターゲットの移動)
ダブルDQNQ値の過大評価
DQNとの決闘より良い価値推定
優先再生重要な経験のサンプルをもっと見る
レインボーすべての改善点を組み合わせる