Aller au contenu principal

Module 4 — Gradients qui disparaissent et qui explosent

Le module 3 a montré que BPTT décompose la dérivée en une somme de termes, dont chacun contient un produit de jacobiennes sur les pas séparant tt de TT. Ce produit n'est pas anodin : il transforme un facteur légèrement inférieur ou supérieur à 1 en 0 ou en \infty dès qu'il est répété une centaine de fois. Voici pourquoi, et ce qu'on y peut sans changer d'architecture.

Le mécanisme en une équation

Le gradient qui remonte de hTh_T vers hth_t traverse TtT - t pas, et à chaque pas il est multiplié par la jacobienne locale :

hTht=k=t+1Thkhk1\frac{\partial h_T}{\partial h_t} = \prod_{k=t+1}^{T} \frac{\partial h_k}{\partial h_{k-1}}

Pour un RNN simple avec activation tanh\tanh, cette jacobienne s'écrit diag(tanh(zk))Wh\operatorname{diag}(\tanh'(z_k)) \cdot W_h, où zkz_k est l'entrée préactivation. La dérivée tanh\tanh' est comprise dans [0,1][0, 1], donc essentiellement inférieure à 1 sauf autour de zéro. La norme du produit se comporte alors comme la puissance (Whσ)Tt(\|W_h\| \cdot \sigma)^{T - t}, avec σ\sigma un facteur moyen inférieur à 1.

Deux régimes se dégagent :

  • Si Whσ<1\|W_h\| \cdot \sigma < 1, la norme du gradient tend vers 0 de manière exponentielle : le modèle n'apprend plus les dépendances longues.
  • Si Whσ>1\|W_h\| \cdot \sigma > 1, la norme explose : les mises à jour deviennent gigantesques et la perte diverge en NaN.

Démonstration numérique sur 50 pas

Voici une expérience minimale, sans réseau, qui montre le phénomène pour différents choix de WhW_h.

import numpy as np
import matplotlib.pyplot as plt

def evolution_norme(spectre, T=50, H=64, essais=20):
"""Retourne la norme moyenne de la jacobienne cumulee apres T pas."""
normes = []
for _ in range(essais):
# W_h avec rayon spectral fixe
M = np.random.randn(H, H)
U, _, Vt = np.linalg.svd(M, full_matrices=False)
W = spectre * (U @ Vt) # orthogonal, echelle = spectre
produit = np.eye(H)
trace = []
for _ in range(T):
produit = W @ produit
trace.append(np.linalg.norm(produit))
normes.append(trace)
return np.mean(normes, axis=0)

for spectre in [0.6, 0.9, 1.0, 1.1, 1.4]:
plt.plot(evolution_norme(spectre), label=f"spectre = {spectre}")
plt.yscale("log")
plt.xlabel("pas de temps")
plt.ylabel("norme cumulee")
plt.legend()
plt.title("Produit de 50 jacobiennes selon le rayon spectral")
plt.show()

Le tracé montre trois régimes : décroissance exponentielle vers 0 pour 0,60{,}6 et 0,90{,}9, plateau autour de 1 pour 1,01{,}0, croissance exponentielle explosive pour 1,11{,}1 et 1,41{,}4. La frontière est très étroite : entre 0,9 et 1,1, l'apprentissage passe de « impossible » à « instable ».

L'écrêtage du gradient : le pansement de l'explosion

L'explosion se traite bien, la disparition ne se traite pas. L'écrêtage consiste à borner la norme du gradient avant la mise à jour :

ggmin(1,θg)g \leftarrow g \cdot \min\left(1, \frac{\theta}{\|g\|}\right)

θ\theta est un seuil (typiquement 1,0 ou 5,0). Si la norme du gradient dépasse θ\theta, on la redescend à θ\theta sans changer sa direction.

En Keras, c'est un argument de l'optimiseur :

from tensorflow import keras

# clipnorm ecrete la norme globale du gradient sur tous les parametres
optimiseur = keras.optimizers.Adam(learning_rate=1e-3, clipnorm=1.0)
modele.compile(optimizer=optimiseur, loss="mse")

En PyTorch, on l'appelle explicitement entre backward et step :

import torch

for lot in loader:
optimiseur.zero_grad()
perte = fonction_perte(modele(lot["x"]), lot["y"])
perte.backward()
torch.nn.utils.clip_grad_norm_(modele.parameters(), max_norm=1.0)
optimiseur.step()

L'écrêtage sauve un entraînement du NaN mais ne rend pas la mémoire longue apprenable : il empêche la divergence, il ne guérit pas la disparition. Pour la disparition, la vraie réponse est architecturale (LSTM au module 5, GRU au module 6).

L'initialisation orthogonale : le bon point de départ

Une matrice orthogonale a toutes ses valeurs singulières égales à 1, ce qui fait exactement le régime marginal Wh=1\|W_h\| = 1. C'est pourquoi les couches récurrentes en Keras utilisent par défaut orthogonal pour la matrice WhW_h et glorot_uniform pour WxW_x.

from tensorflow.keras import layers, initializers

cellule = layers.SimpleRNN(
64,
kernel_initializer=initializers.GlorotUniform(), # W_x
recurrent_initializer=initializers.Orthogonal(), # W_h
)

Ne pas changer ces valeurs par défaut est le conseil le plus rentable de ce module. Une initialisation gaussienne standard sur WhW_h démarre avec un rayon spectral >1> 1 et explose au premier pas de gradient.

Diagnostiquer, pas deviner

Trois signaux dans les journaux d'entraînement indiquent lequel des deux maux vous frappe.

La perte devient NaN ou inf en quelques itérations. C'est l'explosion. Diagnostic immédiat : activer clipnorm=1.0. Si le problème persiste, réduire le taux d'apprentissage d'un facteur 10.

La perte diminue régulièrement mais n'apprend pas les dépendances au-delà d'une dizaine de pas. C'est la disparition. Vérifier avec l'expérience « deux longueurs » du module 3 : entraîner sur T=10T = 10 et sur T=100T = 100. Si les scores sont identiques, tout ce qui est au-delà de 10 est ignoré.

Le gradient a une norme qui s'écroule à 0 sur les paramètres profonds du graphe déplié. Enregistrable via TensorBoard ou un journal manuel. C'est le symptôme direct.

import tensorflow as tf

@tf.function
def etape_avec_norme(x, y):
with tf.GradientTape() as ruban:
perte = fonction_perte(modele(x), y)
grads = ruban.gradient(perte, modele.trainable_variables)
norme = tf.linalg.global_norm(grads)
return perte, norme

for lot_x, lot_y in jeu:
perte, norme = etape_avec_norme(lot_x, lot_y)
tf.summary.scalar("norme_gradient", norme, step=etape_actuelle)

Un graphe de norme_gradient qui descend sous 10410^{-4} et y reste, alors que la perte plafonne, signe la disparition.

Ce qui ne marche pas

Trois idées séduisantes n'apportent rien.

Augmenter le taux d'apprentissage pour « compenser » un gradient qui disparaît accélère la divergence des paramètres qui reçoivent encore du gradient (les derniers pas), sans réveiller les autres.

Remplacer tanh\tanh par ReLU dans un RNN simple élimine la borne supérieure de tanh\tanh' et rend l'explosion encore plus facile. C'est l'inverse de ce qu'on cherche.

Ajouter du dropout sur WhW_h dégrade encore davantage la propagation du gradient au lieu de la stabiliser. Le dropout récurrent propre (même masque à tous les pas) fonctionne, mais il est vu au module 7.

L'écrêtage n'est pas un substitut à un LSTM

Un dépannage d'urgence sur un RNN simple qui explose n'est pas la même chose qu'une solution à long terme. Si votre tâche a des dépendances au-delà de 20 à 30 pas, la disparition frappera indépendamment de l'écrêtage. Passez à un LSTM ou un GRU dès que le diagnostic pointe la disparition ; l'écrêtage seul ne l'atteindra pas.

Un seul chiffre à surveiller : la norme globale du gradient

La courbe global_norm en fonction des itérations est un signal extraordinairement dense. Une pointe brutale annonce une divergence à venir. Une courbe qui s'écrase à 0 annonce la disparition. Une courbe stable autour d'une valeur modérée est le signe d'un entraînement sain. Ajoutez-la systématiquement à vos tableaux de bord.

En résumé

  • Le gradient qui remonte TtT - t pas traverse un produit de jacobiennes dont la norme se comporte comme (Whσ)Tt(\|W_h\| \cdot \sigma)^{T - t} : il disparaît ou explose selon que ce facteur est inférieur ou supérieur à 1.
  • L'écrêtage (clipnorm) traite l'explosion en bornant la norme du gradient sans changer sa direction ; il ne guérit pas la disparition.
  • L'initialisation orthogonale de WhW_h (par défaut en Keras) place le point de départ à la frontière stable ; ne pas la changer sans raison.
  • Diagnostiquer avec la norme globale du gradient : NaN rapide signe l'explosion, écrasement à 0 signe la disparition, et les deux appellent des réponses différentes — l'une algorithmique, l'autre architecturale.

Module suivant : le LSTM répond à la disparition en introduisant un chemin dédié où le gradient circule sans multiplication répétée.