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 de . Ce produit n'est pas anodin : il transforme un facteur légèrement inférieur ou supérieur à 1 en 0 ou en 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 vers traverse pas, et à chaque pas il est multiplié par la jacobienne locale :
Pour un RNN simple avec activation , cette jacobienne s'écrit , où est l'entrée préactivation. La dérivée est comprise dans , donc essentiellement inférieure à 1 sauf autour de zéro. La norme du produit se comporte alors comme la puissance , avec un facteur moyen inférieur à 1.
Deux régimes se dégagent :
- Si , la norme du gradient tend vers 0 de manière exponentielle : le modèle n'apprend plus les dépendances longues.
- Si , 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 .
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 et , plateau autour de 1 pour , croissance exponentielle explosive pour et . 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 :
où est un seuil (typiquement 1,0 ou 5,0). Si la norme du gradient dépasse , on la redescend à 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 . C'est pourquoi les couches
récurrentes en Keras utilisent par défaut orthogonal pour la matrice
et glorot_uniform pour .
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 démarre avec un rayon spectral 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 et sur . 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 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 par ReLU dans un RNN simple élimine la borne supérieure de et rend l'explosion encore plus facile. C'est l'inverse de ce qu'on cherche.
Ajouter du dropout sur 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.
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.
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 pas traverse un produit de jacobiennes dont la norme se comporte comme : 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 (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 :
NaNrapide 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.