Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

Produits jacobien-vecteur en mode inverse (VJP)

Open In Colab

Produits jacobien-vecteur en mode inverse (VJP)

Ce carnet illustre le fonctionnement de la différentiation automatique en mode inverse (reverse mode). Nous verrons que :

  • la jacobienne d’une composition se factorise en un produit de jacobiennes par couche,

  • extraire une ligne de cette jacobienne revient à multiplier à gauche par un vecteur de base,

  • les VJP (vector-Jacobian products) effectuent cette extraction sans matérialiser les jacobiennes,

  • il faut autant de passes arrière qu’il y a de sorties.

import jax
import jax.numpy as jnp

key = jax.random.PRNGKey(0)

def sigmoid(x):
    return 0.5 * (1 + jnp.tanh(x / 2))

def affine(W, x):
    """Couche : transformation linéaire suivie d'une sigmoïde."""
    return sigmoid(W @ x)

# Entrée et poids aléatoires
x = jax.random.normal(key, (4,))
W1 = jax.random.normal(key, (2, 4))
W2 = jax.random.normal(key, (6, 2))
W3 = jax.random.normal(key, (2, 6))
W0409 22:49:19.527717 39649308 cpp_gen_intrinsics.cc:74] Empty bitcode string provided for eigen. Optimizations relying on this IR will be disabled.

Passe avant

Le réseau est une composition de trois couches σ(Wk ⋅ )\sigma(W_k \,\cdot\,) :

f(x)=σ ⁣(W3  σ ⁣(W2  σ(W1x)))f(\mathbf{x}) = \sigma\!\big(W_3 \;\sigma\!\big(W_2 \;\sigma(W_1 \mathbf{x})\big)\big)

Calculons les activations intermédiaires v1,v2,v3\mathbf{v}_1, \mathbf{v}_2, \mathbf{v}_3 :

v1 = affine(W1, x)   # R^4 → R^2
v2 = affine(W2, v1)  # R^2 → R^6
v3 = affine(W3, v2)  # R^6 → R^2

print(f"x  = {x}")
print(f"v1 = {v1}")
print(f"v2 = {v2}")
print(f"v3 = {v3}  ← sortie f(x)")
x  = [ 1.6226422   2.0252647  -0.43359444 -0.07861735]
v1 = [0.9990218  0.18136695]
v2 = [0.8795707  0.38997227 0.49990344 0.40007555 0.6204294  0.8609016 ]
v3 = [0.77577215 0.34936523]  ← sortie f(x)

Jacobienne par la règle de chaîne

La jacobienne de ff en x\mathbf{x} se factorise en un produit de jacobiennes par couche :

Jf(x)=J3  J2  J1J_f(\mathbf{x}) = J_3 \; J_2 \; J_1

où JkJ_k est la jacobienne de la kk-ième couche par rapport à son entrée, évaluée à l’activation intermédiaire correspondante. Vérifions que ce produit donne le même résultat que jax.jacobian appliqué directement à ff :

jac_affine = jax.jacobian(affine, argnums=1)

J1 = jac_affine(W1, x)
J2 = jac_affine(W2, v1)
J3 = jac_affine(W3, v2)

# Jacobienne complète calculée directement
f = lambda x: affine(W3, affine(W2, affine(W1, x)))
J_direct = jax.jacobian(f)(x)

# Jacobienne par produit des jacobiennes par couche
J_chain = J3 @ J2 @ J1

print("jax.jacobian(f)(x) :")
print(J_direct)
print("\nJ3 @ J2 @ J1 :")
print(J_chain)
jax.jacobian(f)(x) :
[[ 0.00265777 -0.01498167 -0.00759213  0.00759245]
 [-0.00255269  0.01369359  0.00703029 -0.00699869]]

J3 @ J2 @ J1 :
[[ 0.00265777 -0.01498167 -0.00759213  0.00759245]
 [-0.00255269  0.01369359  0.00703029 -0.00699869]]

Extraire les lignes : une passe par sortie

La jacobienne JfJ_f est une matrice 2×42 \times 4 (2 sorties, 4 entrées). Chaque ligne correspond au gradient d’une composante de la sortie par rapport à toutes les entrées. Pour extraire la ii-ème ligne, on multiplie à gauche par le ii-ème vecteur de base ei\mathbf{e}_i :

ei⊤ Jf=ei⊤ J3 J2 J1\mathbf{e}_i^\top \, J_f = \mathbf{e}_i^\top \, J_3 \, J_2 \, J_1

Si ff a mm sorties, il faut mm tels produits pour reconstituer la jacobienne complète.

e = jnp.eye(2)

row_0 = e[0] @ J3 @ J2 @ J1
row_1 = e[1] @ J3 @ J2 @ J1

print(f"e_0 @ J : {row_0}")
print(f"e_1 @ J : {row_1}")
print(f"\nLignes de J_direct :")
print(f"  ligne 0 : {J_direct[0]}")
print(f"  ligne 1 : {J_direct[1]}")
e_0 @ J : [ 0.00265777 -0.01498167 -0.00759213  0.00759245]
e_1 @ J : [-0.00255269  0.01369359  0.00703029 -0.00699869]

Lignes de J_direct :
  ligne 0 : [ 0.00265777 -0.01498167 -0.00759213  0.00759245]
  ligne 1 : [-0.00255269  0.01369359  0.00703029 -0.00699869]

VJP : multiplier sans matérialiser la jacobienne

Le produit ei⊤J3 J2 J1\mathbf{e}_i^\top J_3 \, J_2 \, J_1 se calcule de gauche à droite en multipliant successivement par chaque jacobienne. Mais dans un vrai réseau, ces jacobiennes intermédiaires peuvent être très grandes.

La fonction jax.vjp résout ce problème : elle calcule le produit v⊤Jk\mathbf{v}^\top J_k sans jamais construire JkJ_k explicitement. Lors de la passe avant, elle enregistre les quantités nécessaires, puis retourne une fonction qui accepte un vecteur cotangent v\mathbf{v} et produit v⊤Jk\mathbf{v}^\top J_k.

Vérifions sur la dernière couche que le VJP donne le même résultat que e0⊤J3\mathbf{e}_0^\top J_3 :

_, vjp_layer3 = jax.vjp(affine, W3, v2)

# vjp_layer3(v) retourne (v @ daffine/dW3, v @ daffine/dv2)
# Le deuxième élément est v @ J3
vjp_result = vjp_layer3(e[0])[1]
direct_result = e[0] @ J3

print(f"VJP :      {vjp_result}")
print(f"e_0 @ J3 : {direct_result}")
VJP :      [ 0.28225815  0.35229424 -0.07542363 -0.01367547  0.03063096 -0.16909465]
e_0 @ J3 : [ 0.28225815  0.35229424 -0.07542363 -0.01367547  0.03063096 -0.16909465]

Chaîner les VJP : la passe arrière

Pour calculer ei⊤J3 J2 J1\mathbf{e}_i^\top J_3 \, J_2 \, J_1 sans matérialiser aucune jacobienne, on enchaîne les VJP de la dernière couche vers la première. Le vecteur adjoint a\mathbf{a} est propagé vers l’arrière :

a3=ei,a2=a3⊤J3,a1=a2⊤J2,a0=a1⊤J1\mathbf{a}_3 = \mathbf{e}_i, \qquad \mathbf{a}_2 = \mathbf{a}_3^\top J_3, \qquad \mathbf{a}_1 = \mathbf{a}_2^\top J_2, \qquad \mathbf{a}_0 = \mathbf{a}_1^\top J_1

Le résultat a0\mathbf{a}_0 est la ii-ème ligne de la jacobienne. C’est exactement le mécanisme de la rétropropagation.

# Passe arrière pour la première sortie (i = 0)
adjoint = e[0]                                       # a_3 = e_0
adjoint = jax.vjp(affine, W3, v2)[1](adjoint)[1]    # a_2 = a_3 @ J3
adjoint = jax.vjp(affine, W2, v1)[1](adjoint)[1]    # a_1 = a_2 @ J2
adjoint = jax.vjp(affine, W1, x)[1](adjoint)[1]     # a_0 = a_1 @ J1

print(f"Passe arrière (VJP chaînés) : {adjoint}")
print(f"Ligne 0 de J_direct :         {J_direct[0]}")
Passe arrière (VJP chaînés) : [ 0.00265777 -0.01498167 -0.00759213  0.00759245]
Ligne 0 de J_direct :         [ 0.00265777 -0.01498167 -0.00759213  0.00759245]

Toute la passe arrière tient en une seule expression, où les appels VJP s’emboîtent de l’intérieur vers l’extérieur :

row_0_vjp = jax.vjp(affine, W1, x)[1](
    jax.vjp(affine, W2, v1)[1](
        jax.vjp(affine, W3, v2)[1](e[0])[1]
    )[1]
)[1]

print(f"VJP chaînés :  {row_0_vjp}")
print(f"J_direct[0] :  {J_direct[0]}")
VJP chaînés :  [ 0.00265777 -0.01498167 -0.00759213  0.00759245]
J_direct[0] :  [ 0.00265777 -0.01498167 -0.00759213  0.00759245]

Récapitulatif

  • La jacobienne d’une composition se factorise : Jf=JL⋯J2 J1J_f = J_L \cdots J_2 \, J_1.

  • Chaque ligne de JfJ_f s’obtient par un produit vecteur-jacobienne, en partant d’un vecteur de base ei\mathbf{e}_i et en propageant de la dernière couche vers la première.

  • jax.vjp calcule ces produits sans matérialiser les jacobiennes intermédiaires : seules les activations de la passe avant sont stockées.

  • Pour une fonction f:Rn→Rmf : \mathbb{R}^n \to \mathbb{R}^m, il faut mm passes arrière pour reconstituer la jacobienne. Quand m=1m = 1 (fonction de perte scalaire), une seule passe arrière suffit pour obtenir le gradient complet par rapport à toutes les entrées et tous les paramètres.