Deep Q-Networks
Be able to train a DQN with a replay buffer and a target network.
Prerequisites
- DNeural networks — the forward pass with matricesrequired
- EQ-learningrequired
Intuition
Q-learning with a table works. Swap the table for a neural network and everything breaks — for three reasons, each with its own solution.
| Problem | Why | Solution |
|---|---|---|
| Correlated data | consecutive steps are nearly identical; SGD assumes independent samples | a replay buffer — store transitions, sample randomly from it |
| A moving target | the target is computed with the same network being trained; chasing your own tail | a target network — a frozen copy updated now and then |
| Overestimation | a over noisy estimates is systematically too high | Double DQN — pick the action with one network, value it with the other |
It was the combination of the first two that made DQN play Atari in 2013 — not the network architecture, which was entirely ordinary.
The overestimation is subtle and deserves its own sentence: even if every Q estimate is unbiased, the expectation of the maximum is larger than the maximum of the expectations. The noise always goes the same way when you take a max.
Code
import random, torch, torch.nn as nn
from collections import deque
class Replay:
def __init__(self, capacity=100_000):
self.buf = deque(maxlen=capacity) # a ring buffer: the oldest falls out automatically
def add(self, s, a, r, s2, done):
self.buf.append((s, a, r, s2, float(done)))
def sample(self, n, rng=random):
batch = rng.sample(self.buf, n)
return [torch.as_tensor(x, dtype=torch.float32) for x in zip(*batch)]
def __len__(self):
return len(self.buf)
q = nn.Sequential(nn.Linear(4, 128), nn.ReLU(), nn.Linear(128, 2))
target_q = nn.Sequential(nn.Linear(4, 128), nn.ReLU(), nn.Linear(128, 2))
target_q.load_state_dict(q.state_dict())
opt = torch.optim.Adam(q.parameters(), lr=1e-3)
def train_step(replay, batch=64, gamma=0.99, double=True):
s, a, r, s2, done = replay.sample(batch)
qsa = q(s).gather(1, a.long().unsqueeze(1)).squeeze(1)
with torch.no_grad():
if double:
b = q(s2).argmax(1, keepdim=True) # CHOOSE with the online network
q2 = target_q(s2).gather(1, b).squeeze(1) # VALUE with the target network
else:
q2 = target_q(s2).max(1).values
target = r + gamma * q2 * (1 - done) # (1 - done) — the classic bug
loss = nn.functional.smooth_l1_loss(qsa, target) # Huber copes with outliers better than MSE
opt.zero_grad(); loss.backward()
nn.utils.clip_grad_norm_(q.parameters(), 10.0)
opt.step()
return float(loss)
# The target network is updated rarely — every 1 000 steps, or softly every step:
def soft_update(tau=0.005):
for p, pt in zip(q.parameters(), target_q.parameters()):
pt.data.mul_(1 - tau).add_(tau * p.data)
(1 - done) is the bug that costs the most hours. Without it the agent bootstraps on the value of a state beyond the end of the episode. The Q values grow, the loss looks plausible, and the agent never learns. It shows up in no error message.
The debugging order when a DQN does not learn:
- Check
(1 - done). - Print the mean Q over time — if it grows without bound something is wrong with the target.
- Check that
epsilonreally is decreasing. - Check that the buffer has had time to fill before training starts.
- Only then: the hyperparameters.
Mastery means
- Explains why the replay buffer and the target network are needed
- Implements the DQN target correctly
- Knows about the overestimation problem and Double DQN
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Playing Atari with Deep Reinforcement Learning — arXiv (open access; licence per article)
- arXiv — Deep Reinforcement Learning with Double Q-learning — arXiv (open access; licence per article)
- OpenAI Spinning Up in Deep RL (MIT) — MIT