Différentiation automatique
PyTorch Autograd Explained - In-depth Tutorial - YouTube 2.5. Automatic Differentiation — Dive into Deep Learning documentation
Différentiation automatique avec JAX
JAX est une bibliothèque conçue pour le calcul numérique et scientifique haute performance, qui exploite les capacités du matériel moderne. Elle étend NumPy et permet la différentiation automatique, donc le calcul efficace des gradients, essentiels aux problèmes d’optimisation et aux algorithmes d’apprentissage automatique.
Différentiation automatique (autodiff)
La différentiation automatique évalue les dérivées en composant les règles des opérations élémentaires du programme par la règle de la chaîne. Contrairement au calcul symbolique, elle n'a pas besoin de construire une formule simplifiée ; contrairement aux différences finies, elle n'estime pas la dérivée à partir de valeurs voisines. Elle dérive le calcul implémenté, avec les erreurs d'arrondi et les règles disponibles. Une discontinuité, une décision entière, une opération non prise en charge ou une dérivée inexistante ne deviennent pas dérivables grâce à l'autodiff. Les dérivées supérieures exigent une différentiabilité suffisante et peuvent coûter cher. Revoyez les dérivées et les dimensions des applications linéaires avant les produits jacobiens.
Autodiff en mode avant et inverse
Pour , le jacobien a pour coefficients et pour taille . Le mode avant propage une direction d'entrée vers . Le mode inverse propage des poids de sortie vers . Des vecteurs de base comme graines donnent des colonnes ou lignes individuelles ; les graines ne se limitent pas à une coordonnée.
La transposée vient de la règle de chaîne multivariée, pas d’une convention propre à JAX. Soit et une perte scalaire . Chaque entrée peut agir sur la perte par chaque sortie, d’où
Le mode inverse prend donc . La somme explique aussi pourquoi les contributions de chemins distincts s’accumulent : la règle de chaîne scalaire s’applique à toutes les coordonnées intermédiaires.
Un jacobien complet demande donc environ parcours avant ou parcours inverses. Le mode avant convient à peu d'entrées et beaucoup de sorties ; le mode inverse à beaucoup d'entrées et peu de sorties, surtout une perte scalaire. Il fonde la rétropropagation, mais doit conserver ou recalculer les intermédiaires pour le parcours arrière. Ce sont des compromis entre travail et mémoire, pas un classement universel de vitesse.
Graphe de différentiation
Un graphe de différentiation est une représentation graphique du calcul d’une fonction qui facilite le calcul de dérivées. Il représente la suite des opérations, les variables intermédiaires et leurs dépendances. Les algorithmes de différentiation automatique l’emploient pour appliquer la règle de la chaîne efficacement et précisément.
Comment JAX implémente l’autodiff
JAX utilise les modes avant et inverse pour calculer efficacement gradients, jacobiens et hessiens. Il fournit notamment :
grad: calcule le gradient d’une fonction. Il emploie le mode inverse, particulièrement utile lorsque la dimension d’entrée est très supérieure à celle de sortie.jax.jvp: calcule un produit jacobien-vecteur, l’opération élémentaire du mode avant.jax.vjp: calcule un produit vecteur-jacobien, l’opération élémentaire du mode inverse.
Exemple d’utilisation
Voici comment calculer avec JAX le gradient d’une fonction :
import jax
import jax.numpy as jnp
def f(x):
return jnp.sin(x) * jnp.cos(x)
grad_f = jax.grad(f)
print(grad_f(1.0)) # Computes the gradient of f at x = 1.0
Ce code calcule la dérivée de au point .
Suivre une entrée partagée
Posons , et . En , la graine avant donne et , donc . C'est .
En mode inverse, on initialise . L'addition transmet . La multiplication apporte à et à ; le sinus apporte à . Il faut additionner les contributions à une même entrée : . Les deux chemins passant par expliquent pourquoi écraser son gradient serait incorrect.
Dans l'exemple JAX précédent, la dérivée attendue est . jax.grad exige une sortie scalaire (forme (), pas un vecteur de longueur un) et dérive habituellement des entrées flottantes. Pour des sorties vectorielles, choisissez jax.jvp, jax.vjp, jax.jacfwd ou jax.jacrev selon le produit ou la matrice recherchés. En un point anguleux, comme zéro pour , la règle choisie par une implémentation ne prouve pas l'existence d'une dérivée classique.