Aller au contenu principal

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 Q[s][a]Q[s][a] n'a plus de sens si ss est un vecteur de flottants. Le Deep Q-Network, ou DQN, remplace la table par un réseau de neurones Qθ(s,a)Q_{\theta}(s, a) 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 10810^8 à 101010^{10} é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 Qθ(s,a)Q_{\theta}(s, a) 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 r+γmaxaQθ(s,a)r + \gamma \max_{a'} Q_{\theta}(s', a') dépend des poids θ\theta que l'entraînement modifie. Ajuster Qθ(s,a)Q_{\theta}(s, a) 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 QQ, cette surestimation s'amplifie à travers les mises à jour successives.

Les deux astuces de DQN

Rejeu d'expérience. Stocker chaque transition (s,a,r,s,fin)(s, a, r, s', \text{fin}) 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 QθQ_{\theta} est celui qu'on entraîne. Le réseau cible QθQ_{\theta^-} est une copie figée qui calcule la cible :

y=r+γmaxaQθ(s,a)y = r + \gamma \max_{a'} Q_{\theta^-}(s', a')

Les poids θ\theta^- sont recopiés depuis θ\theta tous les CC pas (par exemple 1000), ou par moyenne mobile θτθ+(1τ)θ\theta^- \leftarrow \tau \theta + (1 - \tau)\theta^- avec τ\tau 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 :

y=r+γQθ ⁣(s, argmaxaQθ(s,a))y = r + \gamma\, Q_{\theta^-}\!\left(s',\ \arg\max_{a'} Q_{\theta}(s', a')\right)

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'ε\varepsilon 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é.

Diagnostiquer un DQN qui n'apprend pas

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 QθQ_{\theta} 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 CC 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.