🎤 PyTorch-karaoké, ou comment ne plus oublier optimizer.zero_grad()

Beaucoup de débutants en DL se cassent les dents non pas sur l'architecture complexe des réseaux, mais sur le banal boilerplate de PyTorch. L'ordre des appels dans la boucle d'entraînement est ce qu'on doit googler encore et encore jusqu'à ce que cela s'imprime dans le subconscient.

Tu as oublié de mettre le modèle en mode train() ? Tu obtiens des poids incorrects à cause de Dropout/BatchNorm.
Tu as oublié zero_grad() ? Félicitations, les gradients s'accumulent, l'entraînement part en vrille.
Tu as mis step() avant backward() ? Bon, tu vois le truc.

Je suis tombé sur une vidéo qui provoque deux sentiments à la fois : un immense malaise et du respect.
Un gars en slip avec un micro au milieu du bazar a simplement mis les 5 étapes de la rétropropagation sur une mélodie qui ne sort pas de la tête.
Comme si ça marchait mieux que 10 heures de cours ennuyeux donnés par des experts qui lisent des diapos.

Pour ceux qui sont dans le flou, je rappelle l'ordre unique et correct des opérations, à tatouer dans votre subconscient (ou à apprendre la chanson) :

1️⃣ model.train() — on met le modèle en mode combat.
2️⃣ y_pred = model(x) — Forward pass.
3️⃣ loss = loss_fn(y_pred, y) — On calcule à quel point on s'est trompé.
4️⃣ optimizer.zero_grad() — On remet les gradients à zéro avant la nouvelle étape, sinon ils s'accumulent et vous partez dans l'espace.
5️⃣ loss.backward() — On calcule les gradients (rétropropagation).
6️⃣ optimizer.step() — On fait un pas avec l'optimiseur.

Il faudrait aussi une version pour TensorFlow, mais là, j'ai peur qu'il faille écrire un opéra en trois actes juste pour initialiser les variables 🌚

#pour_le_plaisir