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

Lesson 5: DQN — Deep Q-Network & Experience Replay

Function approximation with neural networks. DQN architecture. Experience Replay Buffer. Target network. Double DQN, Dueling DQN, Rainbow DQN.

🧠 AI & ML — Lesson 4 Lesson 5: DQN — Deep Q-Network & Experience Replay

Reinforcement Learning: From Basics to Advanced

Part 2: Deep Reinforcement Learning — Neural Networks meet RL

xdev.asia

Introduction

DQN (Deep Q-Network) — a giant leap when DeepMind used a neural network instead of Q-table, achieving superhuman performance on Atari games (Nature, 2015).


1. From Q-Table to DQN

Q-TableDQN
Store Q(s,a) in tableNeural network approximates Q(s,a;θ)
Only works with discrete, small state spaceWorks with continuous, high-dimensional states
ExactApproximate

2. DQN Architecture

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. Experience Replay

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 Training Loop

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 Improvements

Double DQN

Decouple action selection and evaluation → reduce overestimation:

# 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()

Dueling DQN

Separate value V(s) and advantage 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)

Summary

ImprovementsProblem Solving
Experience ReplayCorrelation between consecutive samples
Target NetworkTraining instability (moving target)
Double DQNQ-value overestimation
Dueling DQNBetter value estimation
Prioritized ReplaySample important experiences more
RainbowCombine all improvements