簡介
DQN(深度 Q 網路) — 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 | 更好的價值評估 |
| 優先重播 | 更多重要經歷範例 |
| 彩虹 | 結合所有改進 |