Comment fonctionne la rétropropagation ?

Questions d’entrevue Deep learning

Intermédiaireretropropagationgradientfondamentaux

La rétropropagation est l’application de la règle de dérivation en chaîne, parcourue de la sortie vers l’entrée. Elle calcule la dérivée de la perte par rapport à chaque paramètre du réseau — c’est-à-dire dans quelle direction et de combien modifier chaque poids pour réduire l’erreur.

Le cycle d’entraînement, en quatre temps :

  1. Passage avant : on calcule la prédiction et la perte, en gardant en mémoire les valeurs intermédiaires de chaque couche.
  2. Passage arrière : on part de la dérivée de la perte par rapport à la sortie, et on la propage couche par couche en la multipliant par les dérivées locales.
  3. Mise à jour : chaque paramètre se déplace dans le sens opposé à son gradient, d’un pas fixé par le taux d’apprentissage.
  4. On recommence sur le lot suivant.

Ce qui fait sa valeur, et qu’on attend en réponse

Ce n’est pas l’idée de dériver, c’est l’efficacité. Un réseau moderne compte des millions de paramètres ; estimer chaque gradient séparément par différences finies demanderait un passage avant complet par paramètre. La rétropropagation obtient tous les gradients en un seul passage arrière, pour un coût du même ordre que le passage avant.

La raison est la réutilisation : le gradient d’une couche se calcule à partir de celui de la couche suivante. On ne recalcule jamais deux fois la même quantité — c’est de la programmation dynamique appliquée à un graphe de calcul.

Les conséquences pratiques à savoir énoncer

Les valeurs intermédiaires doivent être conservées entre les deux passages, puisque les dérivées locales en dépendent. C’est ce qui explique que la mémoire consommée croisse avec la profondeur et avec la taille du lot — et pourquoi une saturation mémoire survient au passage arrière plutôt qu’au passage avant. Le gradient checkpointing échange précisément du calcul contre de la mémoire en réoubliant certaines activations pour les recalculer.

Les gradients se multiplient en remontant. Une chaîne de facteurs inférieurs à 1 tend vers zéro — c’est la disparition du gradient — et une chaîne de facteurs supérieurs à 1 explose. Toute la boîte à outils du domaine découle de là : ReLU, initialisation soignée, normalisation par lots, connexions résiduelles, écrêtage du gradient.

Il faut remettre les gradients à zéro entre deux lots, faute de quoi ils s’accumulent. C’est une erreur classique en PyTorch, où l’accumulation est le comportement par défaut — et un choix délibéré, puisqu’elle permet de simuler un grand lot sur une petite carte graphique.

Et la nuance qui situe la réponse : la rétropropagation calcule les gradients, elle n’optimise pas. C’est l’optimiseur — SGD, Adam — qui décide quoi faire de ces gradients. Confondre les deux est fréquent, et les distinguer clairement fait bonne impression.

Toutes les questions Deep learning