AdamW est un algorithme d'optimisation utilisé en apprentissage automatique et en apprentissage profond pour entraîner des réseaux neuronaux. Il s'agit d'une variante de l'optimiseur Adam qui dissocie la décroissance du poids (weight decay) des mises à jour adaptatives du taux d'apprentissage. L'algorithme a été introduit en 2017 par Ilya Loshchilov et Frank Hutter, et il est devenu un choix standard pour l'entraînement de grands modèles, y compris les transformeurs et les grands modèles de langage. En séparant la décroissance du poids des mises à jour de paramètres basées sur le gradient, AdamW corrige un problème connu de l'optimiseur Adam original, où la décroissance du poids était appliquée d'une manière qui interférait avec les taux d'apprentissage adaptatifs, conduisant à une convergence et une généralisation sous-optimales.
La motivation principale d'AdamW est venue de l'observation que dans Adam, le terme de régularisation L2 (souvent utilisé comme décroissance du poids) est divisé par la racine carrée de la moyenne mobile exponentielle des gradients au carré. Ce couplage fait varier la décroissance effective du poids selon les paramètres et au fil du temps, ce qui peut entraver l'optimisation. AdamW propose une correction simple : appliquer la décroissance du poids directement aux paramètres après la mise à jour du gradient, indépendamment de la mise à l'échelle adaptative du taux d'apprentissage. Cette dissociation a montré qu'elle améliore les performances d'entraînement sur diverses tâches, notamment la classification d'images et la modélisation du langage.
Contexte : Optimisation en apprentissage automatique
L'entraînement d'un réseau neuronal implique de minimiser une fonction de perte, généralement en utilisant des variantes de la descente de gradient stochastique (SGD). Dans SGD, les paramètres du modèle sont mis à jour de manière itérative en se déplaçant dans la direction du gradient négatif de la perte, calculé sur un sous-ensemble aléatoire des données d'entraînement. La taille du pas, ou taux d'apprentissage, contrôle l'ampleur de chaque mise à jour. Au fil des ans, de nombreuses améliorations ont été proposées pour accélérer la convergence et améliorer les performances finales, telles que l'élan (momentum), les taux d'apprentissage adaptatifs et la décroissance du poids.
La décroissance du poids est une technique de régularisation qui pénalise les grandes valeurs de paramètres en ajoutant un terme proportionnel à la somme des poids au carré à la fonction de perte. Dans SGD standard, la décroissance du poids est équivalente à la régularisation L2, mais cette équivalence ne tient plus dans les méthodes adaptatives comme Adam. AdamW a été conçu pour rétablir le comportement souhaité de la décroissance du poids dans les optimiseurs adaptatifs.
L'optimiseur Adam
Adam (Adaptive Moment Estimation) a été introduit par Diederik Kingma et Jimmy Ba en 2014. Il maintient des taux d'apprentissage par paramètre en conservant une moyenne décroissante exponentielle des gradients passés (premier moment) et des gradients au carré passés (second moment). La règle de mise à jour d'Adam est :
\[ \theta_{t+1} = \theta_t - \eta \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} \]
où \(\hat{m}_t\) et \(\hat{v}_t\) sont des estimations corrigées du biais des premier et second moments, \(\eta\) est le taux d'apprentissage, et \(\epsilon\) est une petite constante pour la stabilité numérique. Adam est devenu populaire en raison de sa robustesse aux hyperparamètres et de sa convergence rapide, en particulier pour l'entraînement de réseaux profonds.
Cependant, dans l'implémentation originale d'Adam, la décroissance du poids était implémentée comme une régularisation L2, qui ajoute un terme \(\frac{\lambda}{2} \|\theta\|^2\) à la perte. Ce terme est ensuite inclus dans le gradient, et comme Adam normalise le gradient par le second moment, la décroissance effective du poids devient \(\lambda / \sqrt{\hat{v}_t}\). Cela signifie que les paramètres avec de grands gradients reçoivent moins de régularisation, ce qui peut conduire à un surapprentissage et à une mauvaise généralisation.
L'algorithme AdamW
AdamW modifie la règle de mise à jour en supprimant le terme de régularisation L2 du gradient et en appliquant plutôt la décroissance du poids directement aux paramètres après la mise à jour adaptative. La règle de mise à jour devient :
\[ \theta_{t+1} = \theta_t - \eta \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} - \eta \lambda \theta_t \]
où \(\lambda\) est le coefficient de décroissance du poids. Cette dissociation garantit que la décroissance du poids est appliquée uniformément à tous les paramètres, quelle que soit l'ampleur de leurs gradients. Les auteurs ont soutenu que cela conduit à une meilleure généralisation et à un entraînement plus stable, en particulier lors de l'utilisation de taux d'apprentissage élevés.
Dans leur article, Loshchilov et Hutter ont démontré qu'AdamW surpasse Adam avec régularisation L2 sur plusieurs tâches de référence, notamment la classification d'images sur CIFAR-10 et la modélisation du langage sur Penn Treebank. Ils ont également montré qu'AdamW est plus robuste au choix des hyperparamètres de taux d'apprentissage et de décroissance du poids.
Impact sur l'apprentissage profond
AdamW a eu un impact significatif sur le domaine de l'intelligence artificielle. C'est désormais l'optimiseur par défaut dans de nombreux frameworks d'apprentissage profond populaires, tels que PyTorch et TensorFlow, et il est largement utilisé pour entraîner les transformeurs et les grands modèles de langage. Par exemple, de nombreux modèles développés par des organisations comme OpenAI, Anthropic et Google DeepMind utilisent AdamW dans leurs pipelines d'entraînement.
La décroissance du poids dissociée dans AdamW a été particulièrement bénéfique pour l'entraînement de modèles à grande échelle, où la régularisation est cruciale pour éviter le surapprentissage sur des ensembles de données massifs. Il a également été démontré qu'il améliore la vitesse de convergence et les performances finales par rapport à d'autres optimiseurs comme SGD avec élan ou Adam avec régularisation L2.
Comparaison avec d'autres optimiseurs
AdamW est souvent comparé à d'autres optimiseurs adaptatifs tels que AdaGrad, RMSProp et Adam. Alors qu'AdaGrad et RMSProp ajustent les taux d'apprentissage en fonction des gradients historiques, Adam combine à la fois l'élan et les taux d'apprentissage adaptatifs. AdamW améliore Adam en corrigeant le problème de décroissance du poids, ce qui le rend plus fiable pour une large gamme de tâches.
Un autre optimiseur connexe est SGD avec élan, qui est plus simple mais nécessite souvent un réglage minutieux du calendrier du taux d'apprentissage. AdamW offre un bon équilibre entre facilité d'utilisation et performances, c'est pourquoi il est devenu un choix privilégié pour de nombreux praticiens. Cependant, certaines études ont montré que SGD avec élan peut obtenir une meilleure généralisation sur certaines tâches s'il est correctement réglé, mais AdamW reste compétitif et plus robuste aux choix d'hyperparamètres.
Considérations pratiques
Lors de l'utilisation d'AdamW, il y a quelques considérations pratiques. Le coefficient de décroissance du poids \(\lambda\) est généralement défini sur une petite valeur, comme 0,01 ou 0,1, mais il peut nécessiter un réglage pour des tâches spécifiques. Le taux d'apprentissage est souvent défini sur une valeur comme 1e-4 ou 3e-4 pour l'entraînement des transformeurs. De plus, AdamW bénéficie souvent d'un calendrier d'échauffement du taux d'apprentissage, où le taux d'apprentissage augmente progressivement d'une petite valeur à la valeur cible sur les premiers milliers de pas.
En pratique, AdamW est implémenté dans la plupart des bibliothèques d'apprentissage profond, de sorte que les utilisateurs peuvent simplement spécifier l'optimiseur et définir le paramètre de décroissance du poids. Par exemple, dans PyTorch, on peut utiliser torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01). Cette simplicité a contribué à son adoption généralisée.
Extensions et variantes
Plusieurs extensions et variantes d'AdamW ont été proposées. Par exemple, AdamW avec décroissance du poids dissociée a été combiné avec d'autres techniques comme Lookahead, qui maintient un ensemble de paramètres plus lent pour des mises à jour plus fluides. Une autre variante est AdamW avec des calendriers de taux d'apprentissage par recuit cosinusoïdal, qui a montré qu'elle améliore les performances sur les tâches de classification d'images.
Dans le contexte des grands modèles de langage, AdamW est souvent utilisé avec un entraînement en précision mixte et une accumulation de gradient pour gérer de grandes tailles de lots. Certains frameworks offrent également des implémentations fusionnées d'AdamW qui réduisent l'utilisation de la mémoire et améliorent l'efficacité computationnelle, ce qui est important lors de l'entraînement de modèles avec des milliards de paramètres.
Conclusion
AdamW est devenu un outil fondamental dans la boîte à outils de l'apprentissage automatique. En dissociant la décroissance du poids des mises à jour adaptatives du gradient, il corrige un défaut subtil mais important de l'optimiseur Adam original, conduisant à une meilleure généralisation et à un entraînement plus stable. Sa simplicité et son efficacité en ont fait le choix par défaut pour de nombreuses applications d'apprentissage profond, de la classification d'images au traitement du langage naturel. Alors que le domaine continue d'évoluer, AdamW reste une méthode d'optimisation fiable et largement utilisée.
Références
- Loshchilov, I., & Hutter, F. (2017). Decoupled Weight Decay Regularization. arXiv:1711.05101.
- Kingma, D. P., & Ba, J. (2014). Adam: A Method for Stochastic Optimization. arXiv:1412.6980.