Deep Q-networks (DQN)
A DQN replaces the Q-table with a neural network, so an agent can judge situations it has never seen before — plus two tricks that stop it falling apart.
- 15 min read
- 3 reading levels
- Updated
Read these first
On this page 8
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
A deep Q-network is Q-learning with the table swapped for a neural network.
The table had to memorise every situation. The network learns a feel, and applies it to situations it has never met.
The analogy you have already lived
Think about crossing a busy road.
You have never crossed this road, at this time of day, with this exact set of vehicles. Your brain does not hold a memorised entry for it. Yet you judge the gap in about half a second.
You are not looking up a table. You learned a general sense of "a scooter that far away moving that fast" and you apply it to a scene you have never seen.
That is what the network does. The Q-table needed a row for every possible road. The network needs none, because similar-looking situations get similar answers automatically.
Why the table had to go
A table needs one row per situation. That works for a sixteen-square grid.
An Atari screen is 210 by 160 pixels, each with a colour. The number of possible screens is a number with hundreds of digits. There is no computer, and no universe, big enough for that table.
And even if you had it, a table learns nothing from similarity. Move one pixel and the table treats it as an entirely new world it has never seen.
A neural network fixes both problems at once. It takes the situation in, and gives a score for each possible action out.
the situation the network a score for each action
(four numbers, -> (a few dense -> push left: 17.2
or a screen) layers) push right: 21.8
^ pick this oneThe lesson on what a neural network is covers the box in the middle.
And then everything broke
Swapping in a network sounds like a small change. It nearly killed the idea for twenty years, for two reasons.
The lessons come in a chain, and they are all alike. In the last second of play, every frame looks nearly identical. A network trained on a hundred nearly identical examples in a row forgets everything else it knew. Imagine revising for an exam by reading the same paragraph a hundred times.
The target keeps moving. In Q-learning the agent corrects its guess using another of its own guesses. With a table, changing one entry changed one entry. With a network, every update changes every answer a little — including the target you were aiming at. It is like trying to shoot at a target that shifts every time you fire.
Together, these made training diverge. Scores went up, then collapsed to nothing.
The two fixes that made it work
In 2015 a team at DeepMind published two repairs, and the field changed overnight.
A memory of past experience. Instead of learning from what happened one second ago, write everything into a big notebook and train on a random handful pulled from anywhere in it. Old moments and new ones, mixed. This is called a replay buffer, and it breaks the chain of near-identical lessons.
A frozen copy of the network. Keep a second, older copy of the network purely for producing targets. Update it to match the live one every few hundred steps. Now the target sits still long enough to be aimed at. This is called a target network.
playing learning
------- --------
act in the world -> write into the notebook
|
v
pull a random handful
|
v
the live network learns,
aiming at the frozen copy
|
every 250 steps: copy live -> frozenWith those two additions, one network learned forty-nine different Atari games from raw pixels. It reached human level on most of them, with no game-specific changes.
What is honestly hard here
DQN is famous, and it is fragile. Three honest warnings:
It gets worse without warning. An agent scoring well for fifty rounds can collapse to nothing in ten. This is normal, and it is why serious runs save a copy of the best agent so far rather than trusting the final one.
It is enormously wasteful. The original Atari results used 200 million frames per game — about 38 days of nonstop human play, for one game.
It cannot handle a steering wheel. DQN picks from a short list of choices: left, right, fire. A car's steering angle is a smooth range with infinitely many values, and taking the maximum over infinitely many options is not something you can do. That limit is the reason the policy gradient lesson exists.
Remember this
- A DQN replaces the Q-table with a network, so it can generalise to unseen situations.
- Replay buffer: learn from a random mix of old and new experience, never a chain of similar moments.
- Target network: a frozen copy so the thing you aim at stops moving.
What to learn next
- Policy gradients — learning the behaviour directly, with no maximum to take.
- Q-learning — the tabular version this is built on.
- PyTorch basics — the tensor mechanics the code above assumes.
Developer — Code and libraries.
Setup
pip install gymnasium torch numpyCPU-only PyTorch is fine and is what this uses. The script below takes about forty seconds on a laptop. No GPU, no download beyond the packages themselves.
The task
CartPole: a pole balanced on a cart. The observation is four numbers — cart position, cart velocity, pole angle, pole angular velocity. Two actions: push left, push right. Reward is 1 for every step the pole stays up. The episode ends when the pole falls past twelve degrees or after 500 steps.
A score of 500 is the maximum possible.
A complete DQN in seventy lines
import copy
import random
from collections import deque
import gymnasium as gym
import numpy as np
import torch
import torch.nn as nn
torch.manual_seed(0); random.seed(0); np.random.seed(0)
torch.set_num_threads(1)
env = gym.make("CartPole-v1")
env.action_space.seed(0)
def make_net():
return nn.Sequential(nn.Linear(4, 64), nn.ReLU(),
nn.Linear(64, 64), nn.ReLU(),
nn.Linear(64, 2))
online, target = make_net(), make_net()
target.load_state_dict(online.state_dict()) # frozen copy that supplies stable targets
opt = torch.optim.Adam(online.parameters(), lr=1e-3)
buffer = deque(maxlen=10_000) # replay memory: past experience, reused
GAMMA, BATCH, SYNC_EVERY, WARMUP = 0.99, 64, 250, 1000
steps, scores = 0, []
best_avg, best_weights = 0.0, None
for episode in range(300):
state, _ = env.reset(seed=episode)
total, done = 0.0, False
while not done:
epsilon = max(0.05, 1.0 - steps / 5000)
if random.random() < epsilon:
action = env.action_space.sample()
else:
with torch.no_grad():
action = int(online(torch.tensor(state)).argmax())
nxt, reward, terminated, truncated, _ = env.step(action)
buffer.append((state, action, reward, nxt, float(terminated)))
state, total, done = nxt, total + reward, terminated or truncated
steps += 1
if len(buffer) >= WARMUP:
s, a, r, s2, d = zip(*random.sample(buffer, BATCH)) # break the time correlation
s = torch.tensor(np.array(s)); s2 = torch.tensor(np.array(s2))
a = torch.tensor(a)
r = torch.tensor(r, dtype=torch.float32)
d = torch.tensor(d)
q = online(s).gather(1, a.unsqueeze(1)).squeeze(1)
with torch.no_grad():
y = r + GAMMA * target(s2).max(1).values * (1.0 - d)
loss = nn.functional.smooth_l1_loss(q, y)
opt.zero_grad(); loss.backward(); opt.step()
if steps % SYNC_EVERY == 0:
target.load_state_dict(online.state_dict())
scores.append(total)
recent = float(np.mean(scores[-20:]))
if len(scores) >= 20 and recent > best_avg: # RL is not monotonic: keep your best
best_avg, best_weights = recent, copy.deepcopy(online.state_dict())
if (episode + 1) % 50 == 0:
print(f"episode {episode + 1:3d} avg of last 20: {recent:6.1f} epsilon {epsilon:.2f}")
print(f"\nbest 20-episode average seen: {best_avg:.1f}")
online.load_state_dict(best_weights)
evals = []
for k in range(5):
state, _ = env.reset(seed=1000 + k)
total, done = 0.0, False
while not done:
with torch.no_grad():
action = int(online(torch.tensor(state)).argmax())
state, reward, terminated, truncated, _ = env.step(action)
total, done = total + reward, terminated or truncated
evals.append(total)
print("greedy evaluation of the best snapshot:", evals)
env.close()episode 50 avg of last 20: 19.6 epsilon 0.78 episode 100 avg of last 20: 96.3 epsilon 0.26 episode 150 avg of last 20: 87.4 epsilon 0.05 episode 200 avg of last 20: 329.6 epsilon 0.05 episode 250 avg of last 20: 350.4 epsilon 0.05 episode 300 avg of last 20: 347.1 epsilon 0.05 best 20-episode average seen: 475.1 greedy evaluation of the best snapshot: [160.0, 151.0, 163.0, 500.0, 137.0]
Your exact numbers will differ, even with the seeds set, because floating-point summation order varies across PyTorch builds and hardware. The shape is what to check: roughly 20 at the start, a few hundred by episode 200. If yours is flat at 20 after 300 episodes, something is wrong.
The output is honest about DQN's real weakness
Look at episode 100 (96.3), then episode 150 (87.4). It got worse. Then it recovered and passed 300.
Look at the greedy evaluation: [160, 151, 163, 500, 137]. The same weights, on five different starting positions, scored anywhere from 137 to the maximum of 500. That spread is not a bug in the script. It is what DQN on CartPole actually does.
This is why the script keeps best_weights. Take the final network in a real run and you may be taking a snapshot from the middle of a collapse. Nearly every published DQN result is a best-checkpoint result, and papers do not always say so.
Line by line, the parts that matter
target.load_state_dict(online.state_dict()) every 250 steps — the target network. Delete this line, use online for both, and training becomes visibly unstable. It is one line and it is the difference between working and not.
with torch.no_grad(): y = ... — the target must not carry gradients. Without no_grad, gradients flow into the target as well as the prediction, and the network can reduce the loss by moving the target instead of improving the estimate. It will happily do that, and your loss will look excellent while the agent learns nothing.
* (1.0 - d) — zeroes the future for terminated transitions. Note that d stores float(terminated) only, not truncated. When CartPole hits its 500-step limit the pole is still up, and the future genuinely has value. Storing terminated or truncated here is a real bug that quietly caps performance.
random.sample(buffer, BATCH) — the replay buffer. Replace it with the last 64 transitions and watch it fail. Those 64 frames are one moment of one episode, and the network will overfit to it and forget everything else.
smooth_l1_loss — the Huber loss. It behaves like squared error for small errors and like absolute error for large ones, so a single wild TD error cannot produce a gradient that destroys the network. Squared error works on CartPole and is much less forgiving on harder tasks.
nn.Linear(4, 64) — four inputs because CartPole's observation is four numbers, and two outputs because there are two actions. One forward pass gives the score for every action at once, which is why the output layer has one unit per action.
Common mistakes
Training before the buffer has anything in it. The WARMUP check matters. Sampling 64 transitions from a buffer holding 3 gives you the same 3 over and over.
Syncing the target network too often. Set SYNC_EVERY = 1 and you have removed the target network entirely. Set it to 50000 and the targets are so stale that learning crawls. Somewhere between 100 and 10000 steps is normal, and it is worth tuning.
Forgetting torch.no_grad() when choosing an action. It works, but you build a computation graph tens of thousands of times per run for nothing. On CPU this can double your runtime.
Assuming a rising curve means it is finished. Run five seeds. Report the spread. A single rising curve in this field means very little.
Try it yourself
Comment out the if steps % SYNC_EVERY == 0 block so target never updates from online. It will still learn a little at first — the frozen random network gives a stable, wrong target — then stall. Then instead use target = online (same object) and watch the opposite failure: fast early progress followed by a collapse. Seeing both failure modes yourself is worth more than any explanation of why the target network is needed.
What to learn next
- Policy gradients — learning the behaviour directly, with no maximum to take.
- Q-learning — the tabular version this is built on.
- PyTorch basics — the tensor mechanics the code above assumes.
Researcher — Mathematics and papers.
The objective
DQN (Mnih et al., 2015) minimises the squared temporal-difference error with a parameterised action-value function $Q_\theta$:
$$ \mathcal{L}(\theta) = \mathbb{E}{(s,a,r,s') \sim \mathcal{D}} \left[ \left( r + \gamma \max{a'} Q_{\theta^-}(s', a') - Q_\theta(s,a) \right)^2 \right] $$
Symbols: $\mathcal{D}$ the replay buffer, $\theta$ the online parameters, $\theta^-$ the target-network parameters (a periodic copy of $\theta$), $\gamma$ the discount.
The semi-gradient is $\nabla_\theta \mathcal{L} = -\mathbb{E}\left[ \delta \, \nabla_\theta Q_\theta(s,a) \right]$ with $\delta$ the TD error. It is semi-gradient because the dependence of the target on $\theta$ is deliberately ignored. This is not an approximation for convenience: the true gradient of the Bellman residual gives the residual-gradient algorithm (Baird, 1995), which is stable but converges to a different and generally worse solution under function approximation.
The deadly triad
Three ingredients — function approximation, bootstrapping, and off-policy training — can diverge in combination. Any two are safe; all three are not. Baird's counterexample gives divergence on a seven-state MDP with linear features.
DQN uses all three. It has no convergence guarantee, and every component of its design is empirical damage control:
| Component | What it addresses |
|---|---|
| Replay buffer | Non-i.i.d. sequential samples; sample reuse |
| Target network | Non-stationary regression target |
| Reward clipping to $[-1,1]$ | Cross-game gradient scale; changes the objective |
| Huber loss | Gradient blow-up from outlier TD errors |
| Frame stacking (4) | Partial observability — velocity is not in one frame |
Overestimation and Double DQN
$\max_{a'} Q_{\theta^-}(s',a')$ selects and evaluates with the same noisy estimator, giving a positive bias. Van Hasselt, Guez and Silver (2016) showed this bias is large and harmful on Atari, and fixed it by decoupling:
$$ y^{\text{DDQN}} = r + \gamma\, Q_{\theta^-}!\left(s', \arg\max_{a'} Q_{\theta}(s', a')\right) $$
The online network picks; the target network scores. It is a one-line change to the code above and it reliably helps. If you implement one extension, implement this one.
The Rainbow components
Hessel et al. (2018) combined six extensions and measured each by ablation:
- Prioritised replay (Schaul et al., 2016) — sample with probability $\propto |\delta|^\alpha$, correcting the induced bias with importance weights $w_i \propto (1/(N p_i))^\beta$.
- Dueling architecture (Wang et al., 2016) — factor $Q(s,a) = V(s) + A(s,a) - \frac{1}{|\mathcal{A}|}\sum_{a'} A(s,a')$. The subtraction is required for identifiability.
- Multi-step returns — $n$-step targets trade bias for variance; $n = 3$ is the usual choice.
- Distributional RL (Bellemare, Dabney and Munos, 2017) — learn the return distribution over a fixed support rather than its mean, minimising a projected KL.
- Noisy nets (Fortunato et al., 2018) — learned parametric noise in the weights replaces $\epsilon$-greedy with state-dependent exploration.
The ablations found multi-step returns and prioritised replay contributed most; dueling contributed least.
Sample efficiency, the real constraint
The original DQN used 200M frames per game. Subsequent work has compressed this enormously:
- Rainbow — the same score in roughly 40M frames.
- Data-Efficient Rainbow (van Hasselt, Hessel and Aslanides, 2019) — competitive at 100K frames by raising the replay ratio.
- DrQ (Kostrikov et al., 2020) and SPR (Schwarzer et al., 2021) — image augmentation and self-predictive auxiliary losses at 100K frames.
- EfficientZero (Ye et al., 2021) — superhuman Atari at 100K frames using a learned model, which is model-based rather than DQN.
The trend is unambiguous: sample efficiency comes from better representation learning and higher replay ratios, not from better TD targets.
Continuous actions
$\max_{a'} Q(s', a')$ is intractable when $\mathcal{A} \subseteq \mathbb{R}^d$. Three responses:
- Discretise. Suffers exponentially in $d$.
- Learn a deterministic actor to approximate the argmax — DDPG (Lillicrap et al., 2016), TD3 (Fujimoto et al., 2018).
- Learn a stochastic policy directly — see policy gradients and actor-critic.
References
- Mnih et al. (2015), Human-level control through deep reinforcement learning, Nature 518:529-533.
- van Hasselt, Guez and Silver (2016), Deep Reinforcement Learning with Double Q-learning — arxiv.org/abs/1509.06461.
- Schaul et al. (2016), Prioritized Experience Replay — arxiv.org/abs/1511.05952.
- Wang et al. (2016), Dueling Network Architectures for Deep Reinforcement Learning — arxiv.org/abs/1511.06581.
- Hessel et al. (2018), Rainbow: Combining Improvements in Deep Reinforcement Learning — arxiv.org/abs/1710.02298.
- Fujimoto, van Hoof and Meger (2018), Addressing Function Approximation Error in Actor-Critic Methods (TD3) — arxiv.org/abs/1802.09477.
What to learn next
- Policy gradients — learning the behaviour directly, with no maximum to take.
- Q-learning — the tabular version this is built on.
- PyTorch basics — the tensor mechanics the code above assumes.