Aller au contenu principal

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 f:RnRmf:\mathbb R^n\to\mathbb R^m, le jacobien JJ a pour coefficients Jij=fi/xjJ_{ij}=\partial f_i/\partial x_j et pour taille m×nm\times n. Le mode avant propage une direction d'entrée vRnv\in\mathbb R^n vers JvRmJv\in\mathbb R^m. Le mode inverse propage des poids de sortie uRmu\in\mathbb R^m vers JTuRnJ^Tu\in\mathbb R^n. 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 z=f(x)z=f(x) et une perte scalaire L(z)L(z). Chaque entrée peut agir sur la perte par chaque sortie, d’où

Lxj=iLzizixj,xL=JTzL.\frac{\partial L}{\partial x_j}=\sum_i\frac{\partial L}{\partial z_i}\frac{\partial z_i}{\partial x_j}, \qquad \nabla_xL=J^T\nabla_zL.

Le mode inverse prend donc u=zLu=\nabla_zL. 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 nn parcours avant ou mm 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 f(x)=sin(x)cos(x)f(x) = \sin(x) \cos(x) au point x=1.0x = 1.0.

Suivre une entrée partagée

Posons a=xya=xy, b=sinxb=\sin x et f=a+bf=a+b. En (x,y)=(2,3)(x,y)=(2,3), la graine avant (x˙,y˙)=(1,0)(\dot x,\dot y)=(1,0) donne a˙=yx˙+xy˙=3\dot a=y\dot x+x\dot y=3 et b˙=cos2\dot b=\cos2, donc f˙=3+cos22.583853\dot f=3+\cos2\approx2.583853. C'est J(1,0)TJ(1,0)^T.

En mode inverse, on initialise fˉ=1\bar f=1. L'addition transmet aˉ=bˉ=1\bar a=\bar b=1. La multiplication apporte yaˉ=3y\bar a=3 à xˉ\bar x et xaˉ=2x\bar a=2 à yˉ\bar y ; le sinus apporte cos2bˉ\cos2\,\bar b à xˉ\bar x. Il faut additionner les contributions à une même entrée : f=(3+cos2,2)\nabla f=(3+\cos2,2). Les deux chemins passant par xx expliquent pourquoi écraser son gradient serait incorrect.

Dans l'exemple JAX précédent, la dérivée attendue est cos2(1)sin2(1)=cos20.416147\cos^2(1)-\sin^2(1)=\cos2\approx-0.416147. 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 x|x|, la règle choisie par une implémentation ne prouve pas l'existence d'une dérivée classique.

Références et liens utiles

Explorer les liensOuvrir le réseau