Module 7 — Réseaux Q profonds et rejeu d'expérience
Le Q-learning tabulaire s'effondre dès que l'espace d'états devient grand ou continu — la table n'a plus de sens si est un vecteur de flottants. Le Deep Q-Network, ou DQN, remplace la table par un réseau de neurones et applique la même mise à jour. Cela paraît trivial ; c'est en réalité instable, et les deux astuces de DeepMind en 2015 ont été nécessaires pour rendre l'idée exploitable.
Pourquoi la table ne suffit plus
Sur CartPole, l'état est un vecteur de 4 flottants (position, vitesse, angle, vitesse angulaire). Discrétiser en cellules de largeur 0,01 sur chaque dimension donnerait à états, dont l'immense majorité ne serait jamais visitée. La table est à la fois trop grande et trop creuse.
Un réseau apprend une fonction paramétrique qui généralise : deux états proches produisent des valeurs proches sans avoir été rencontrés tous les deux. Cette généralisation est ce qui rend l'espace continu abordable — et c'est aussi ce qui crée les instabilités.
Trois instabilités qui coalisent
Passer du tabulaire au réseau introduit trois problèmes qui interagissent.
Corrélation des échantillons consécutifs. Les transitions successives d'un même épisode se ressemblent énormément (les 4 flottants ne bougent que légèrement à chaque pas). Entraîner un réseau sur ces mini-lots corrélés viole l'hypothèse d'échantillons ind épendants et provoque des oscillations violentes.
Cible qui bouge. La cible TD dépend des poids que l'entraînement modifie. Ajuster vers une cible qui change à chaque pas revient à courir derrière son ombre.
Biais d'optimisme du max. Prendre le maximum sur des estimations bruitées surestime systématiquement la vraie valeur — c'est un résultat classique de statistiques. Sur les , cette surestimation s'amplifie à travers les mises à jour successives.
Les deux astuces de DQN
Rejeu d'expérience. Stocker chaque transition dans un tampon circulaire de taille 100 000 à 1 million. À chaque pas d'entraînement, tirer un mini-lot aléatoire de ce tampon et calculer le gradient. Deux bénéfices : les échantillons du mini-lot sont décorrélés, et chaque transition est réutilisée plusieurs fois — la donnée est chère à collecter en renforcement.
Réseau cible. Maintenir deux copies du réseau. Le réseau en ligne est celui qu'on entraîne. Le réseau cible est une copie figée qui calcule la cible :
Les poids sont recopiés depuis tous les pas (par exemple 1000), ou par moyenne mobile avec petit (0,005). La cible cesse de bouger à chaque pas et l'entraînement devient stable.
Une boucle DQN commentée sur CartPole
import torch, torch.nn as nn, torch.optim as optim
import gymnasium as gym
from collections import deque
import random
env = gym.make("CartPole-v1")
nS, nA = env.observation_space.shape[0], env.action_space.n
reseau = nn.Sequential(nn.Linear(nS, 128), nn.ReLU(), nn.Linear(128, nA))
cible = nn.Sequential(nn.Linear(nS, 128), nn.ReLU(), nn.Linear(128, nA))
cible.load_state_dict(reseau.state_dict())
optimiseur = optim.Adam(reseau.parameters(), lr=1e-3)
tampon = deque(maxlen=50_000)
gamma, taille_lot, epsilon, pas = 0.99, 64, 1.0, 0
for episode in range(500):
s, _ = env.reset(seed=episode)
while True:
pas += 1
if random.random() < epsilon:
a = env.action_space.sample()
else:
with torch.no_grad():
a = int(reseau(torch.tensor(s, dtype=torch.float32)).argmax())
s2, r, termine, tronque, _ = env.step(a)
tampon.append((s, a, r, s2, termine or tronque))
s = s2
if len(tampon) >= taille_lot:
lot = random.sample(tampon, taille_lot)
S = torch.tensor([b[0] for b in lot], dtype=torch.float32)
A = torch.tensor([b[1] for b in lot])
R = torch.tensor([b[2] for b in lot], dtype=torch.float32)
S2 = torch.tensor([b[3] for b in lot], dtype=torch.float32)
F = torch.tensor([b[4] for b in lot], dtype=torch.float32)
with torch.no_grad():
cible_y = R + gamma * cible(S2).max(dim=1).values * (1 - F)
q_pris = reseau(S).gather(1, A.unsqueeze(1)).squeeze(1)
perte = nn.functional.smooth_l1_loss(q_pris, cible_y) # Huber, plus robuste que MSE
optimiseur.zero_grad(); perte.backward()
nn.utils.clip_grad_norm_(reseau.parameters(), 10.0)
optimiseur.step()
if pas % 1000 == 0:
cible.load_state_dict(reseau.state_dict())
if termine or tronque:
break
epsilon = max(0.05, epsilon * 0.995)
Cent lignes qui résolvent CartPole en quelques minutes. Chaque élément — rejeu, réseau cible, perte de Huber, écrêtage du gradient — corrige une instabilité observée expérimentalement.
Double DQN, la correction de 2016
Le max de la cible provoque une surestimation. Double DQN la corrige en séparant la sélection et l'évaluation de l'action :
Le réseau en ligne choisit quelle action serait la meilleure ; le réseau cible évalue cette action. Comme les bruits des deux réseaux ne sont pas parfaitement corrélés, la surestimation se réduit fortement. Sur Atari, Double DQN améliore de nombreux jeux sans coût de calcul notable, et c'est aujourd'hui la variante par défaut.
Ce qui aide en pratique
Perte de Huber (smooth_l1_loss) plutôt que MSE : moins sensible aux transitions à récompense exceptionnelle qui produiraient un gradient énorme sous MSE.
Écrêtage du gradient à norme 10 : borne les mises à jour quand une transition rare provoque une erreur TD géante.
Décroissance douce d' avec un plancher (0,05 à 0,1) : garder un peu d'exploration après convergence évite qu'un événement rare fasse dériver la politique sans être corrigé.
Trois vérifications qui règlent 80 % des cas. Le tampon contient-il suffisamment de transitions avant le premier apprentissage (démarrer après quelques milliers de pas) ? La récompense moyenne progresse-t-elle après quelques centaines d'épisodes, ou reste-t-elle plate — signe d'une exploration insuffisante ou d'un taux d'apprentissage trop grand qui fait diverger le réseau ? Les valeurs Q maximales croissent-elles indéfiniment — signature classique d'une surestimation à corriger par Double DQN ?
En résumé
- Un réseau remplace la table dès que l'espace d'états est grand ou continu ; il généralise entre états proches.
- Trois instabilités menacent : corrélation des échantillons, cible mobile, biais du
max. - Rejeu d'expérience décorrèle et réutilise ; réseau cible fige la cible pendant pas.
- Double DQN réduit la surestimation en séparant sélection et évaluation ; c'est la variante par défaut aujourd'hui.
Module suivant : sortir du cadre des valeurs pour apprendre directement une politique paramétrée.