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 :
Calculons les activations intermédiaires :
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 en se factorise en un produit de jacobiennes par couche :
où est la jacobienne de la -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 à :
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 est une matrice (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 -ème ligne, on multiplie à gauche par le -ème vecteur de base :
Si a sorties, il faut 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 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 sans jamais construire explicitement. Lors de la passe avant, elle enregistre les quantités nécessaires, puis retourne une fonction qui accepte un vecteur cotangent et produit .
Vérifions sur la dernière couche que le VJP donne le même résultat que :
_, 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 sans matérialiser aucune jacobienne, on enchaîne les VJP de la dernière couche vers la première. Le vecteur adjoint est propagé vers l’arrière :
Le résultat est la -è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 : .
Chaque ligne de s’obtient par un produit vecteur-jacobienne, en partant d’un vecteur de base et en propageant de la dernière couche vers la première.
jax.vjpcalcule ces produits sans matérialiser les jacobiennes intermédiaires : seules les activations de la passe avant sont stockées.Pour une fonction , il faut passes arrière pour reconstituer la jacobienne. Quand (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.