IFT3395/IFT6390 — Fondements de l’apprentissage machine
Ce notebook accompagne le Chapitre 7: Réseaux de neurones, sections « Rétropropagation » et « Différentiation automatique ». Vous y dériverez des règles VJP à la main, puis les vérifierez numériquement avec JAX.
Objectifs¶
À la fin de ce TP, vous serez en mesure de:
Interpréter la dérivée totale comme une application linéaire, pas une matrice
Relier les dérivées partielles aux restrictions de (projections et injections)
Appliquer la règle de la chaîne comme composition d’applications linéaires
Calculer des JVP et VJP avec JAX et vérifier par différences finies
Dériver des règles VJP à la main pour des opérations courantes
Composer des VJP pour rétropropager un gradient à travers une chaîne de fonctions
Comparer le coût computationnel du mode avant (JVP) et du mode arrière (VJP)
Prérequis: Ch7, sections « Rétropropagation » et « Différentiation automatique » (en particulier le tableau des règles VJP et l’implémentation minimale avec Var/grad).
Partie 0: Configuration¶
Exécutez cette cellule pour importer les bibliothèques nécessaires. Si vous utilisez Colab, JAX est déjà installé. Localement: pip install jax.
import jax
import jax.numpy as jnp
import numpy as np
import matplotlib.pyplot as plt
# Précision 64 bits — indispensable pour les vérifications par différences finies
jax.config.update("jax_enable_x64", True)
plt.rcParams['figure.figsize'] = (8, 5)
plt.rcParams['font.size'] = 12
print(f"JAX version: {jax.__version__}")
print(f"Plateforme: {jax.devices()[0].platform}")
print("Configuration terminée!")JAX version: 0.9.1
Plateforme: cpu
Configuration terminée!
Partie 1: La notation de Spivak — pourquoi et comment¶
Avant de plonger dans les calculs, clarifions la notation utilisée dans ce TP. Elle diffère peut-être de celle de vos cours de calcul, mais elle a des avantages concrets pour la suite.
Deux traditions de notation¶
La notation de Leibniz — , — est un raccourci commode, mais elle masque un fait important: la dérivée est une fonction (un opérateur), pas un nombre ni un rapport. Quand on écrit , on ne voit pas immédiatement que le résultat dépend du point d’évaluation, ni que la « vraie » dérivée est un objet plus riche qu’une collection de dérivées partielles.
La notation de Spivak adopte une autre perspective: désigne la dérivée de la fonction , évaluée au point . Le résultat est une application linéaire — une fonction qui prend un vecteur et renvoie un vecteur. Pas d’ambiguïté sur ce qu’on dérive ni où.
Pourquoi l’adopter ici¶
Trois raisons pratiques:
Elle correspond à JAX. L’appel
jax.jvp(f, (a,), (v,))applique littéralement l’application linéaire au vecteur . La fonctionfest le premier argument, le point le deuxième, la direction le troisième. Pas de « sortie / entrée » — juste: fonction, point, direction.La règle de la chaîne devient limpide. : c’est une composition de fonctions. Pas de , pas d’indices, pas de conventions d’appariement.
Elle passe à l’échelle. La même écriture fonctionne pour des scalaires, des vecteurs, des matrices ou des tenseurs, sans changer de convention.
Tableau de référence¶
Gardez ce tableau sous la main pour le reste du TP:
| Écriture Spivak | Signification | Écriture Leibniz équivalente |
|---|---|---|
| Dérivée totale de en (application linéaire) | (matrice jacobienne) | |
| Appliquer la dérivée au tangent | (JVP) | |
| Adjoint appliqué au cotangent | (VJP) | |
| Dérivée de la composition |
La perspective opérateur¶
est un opérateur qui transforme une fonction en une nouvelle fonction:
Ensuite, est l’application linéaire obtenue en évaluant au point . Enfin, est cette application linéaire appliquée au vecteur . Trois niveaux d’« application de fonction » — et chacun correspond à un appel JAX:
| Niveau | Notation | JAX |
|---|---|---|
| Opérateur de dérivation | jax.jacobian, jax.jvp, jax.vjp | |
| Évaluation au point | jax.jacobian(f)(a) | |
| Application au vecteur | jax.jvp(f, (a,), (v,))[1] |
Partie 2: La dérivée totale comme application linéaire¶
En calcul à une variable, on écrit : la dérivée est un nombre. En plusieurs variables, la situation est plus riche. Plutôt qu’une matrice de dérivées partielles, la notion fondamentale est celle d’application linéaire.
Définition (Spivak). Soit . La dérivée totale de en est l’unique application linéaire telle que
Trois points à retenir:
est une fonction — elle prend un vecteur et renvoie un vecteur .
La matrice jacobienne représente cette application linéaire dans la base canonique: .
L’écriture est exactement ce que JAX appelle un JVP (Jacobian-vector product).
Travaillons avec un exemple concret tout au long du TP.
def f(x):
"""f: R^2 -> R^3"""
return jnp.array([x[0]**2 + x[1], x[0] * x[1], jnp.sin(x[0])])
a = jnp.array([1.0, 2.0])
print("f(a) =", f(a))f(a) = [3. 2. 0.84147098]
# La matrice jacobienne: la REPRÉSENTATION de Df(a) dans la base canonique
J_f = jax.jacobian(f)(a)
print("Jacobienne J_f(a):")
print(J_f)
print(f"\nTaille: {J_f.shape} (m=3 lignes, n=2 colonnes)")Jacobienne J_f(a):
[[2. 1. ]
[2. 1. ]
[0.54030231 0. ]]
Taille: (3, 2) (m=3 lignes, n=2 colonnes)
Maintenant, appliquons en tant que fonction à un vecteur tangent . On peut le faire de deux façons:
Produit matrice-vecteur: (forme la matrice, puis multiplie)
jax.jvp: calcule directement, sans former la matrice
Les deux donnent le même résultat — mais jax.jvp est plus efficace car il ne matérialise jamais la jacobienne.
v = jnp.array([0.5, -1.0])
# Méthode 1: produit matrice-vecteur (forme la jacobienne complète)
Df_a_v_matrix = J_f @ v
# Méthode 2: jax.jvp — applique Df(a) comme fonction, SANS former la matrice
_, Df_a_v_jvp = jax.jvp(f, (a,), (v,))
print("J_f(a) @ v =", Df_a_v_matrix)
print("jax.jvp(f)(a,v) =", Df_a_v_jvp)
print("Accord:", jnp.allclose(Df_a_v_matrix, Df_a_v_jvp))J_f(a) @ v = [0. 0. 0.27015115]
jax.jvp(f)(a,v) = [0. 0. 0.27015115]
Accord: True
Exercice 1: Vérification de la définition de Spivak ★¶
La définition exige que l’erreur d’approximation linéaire décroisse plus vite que . En posant pour un fixé, le rapport
doit tendre vers 0 quand . Plus précisément, on s’attend à pour une fonction lisse (l’erreur est dominée par le terme quadratique).
Complétez le code ci-dessous pour vérifier cette propriété numériquement, puis tracez en échelle log-log.
t_values = np.logspace(-1, -10, 20)
ratios = []
for t in t_values:
h = t * v
# ============================================
# TODO: Calculez le numérateur (norme de l'erreur d'approximation linéaire)
# et le dénominateur (norme de h), puis le rapport r(t).
#
# numerateur = jnp.linalg.norm(f(a + h) - f(a) - Df(a)(h))
# où Df(a)(h) peut se calculer via J_f @ h ou jax.jvp
# denominateur = jnp.linalg.norm(h)
# ============================================
r_t = None # <- Complétez
ratios.append(r_t)
# Tracé log-log
if ratios[0] is not None:
plt.loglog(t_values, ratios, 'o-', label='$r(t)$')
plt.loglog(t_values, t_values, '--', alpha=0.5, label='pente 1 (référence)')
plt.xlabel('$t$')
plt.ylabel('$r(t)$')
plt.title('Vérification de la définition de Spivak: $r(t)$ doit décroître comme $O(t)$')
plt.legend()
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()
else:
print("Complétez le calcul de r(t)!")Complétez le calcul de r(t)!
Solution Exercice 1 (cliquez pour afficher)
t_values = np.logspace(-1, -10, 20)
ratios = []
for t in t_values:
h = t * v
linear_approx = J_f @ h # ou: jax.jvp(f, (a,), (h,))[1]
numerateur = jnp.linalg.norm(f(a + h) - f(a) - linear_approx)
denominateur = jnp.linalg.norm(h)
ratios.append(float(numerateur / denominateur))
plt.loglog(t_values, ratios, 'o-', label='$r(t)$')
plt.loglog(t_values, t_values, '--', alpha=0.5, label='pente 1 (référence)')
plt.xlabel('$t$')
plt.ylabel('$r(t)$')
plt.title('Vérification de la définition de Spivak: $r(t)$ doit décroître comme $O(t)$')
plt.legend()
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()On observe que r(t) suit la droite de pente 1 en échelle log-log (jusqu’à ce que les erreurs d’arrondi dominent pour de très petites valeurs de t). Cela confirme que le résidu est O(t), donc O(\|\mathbf{h}\|) — l’approximation linéaire est bien meilleure que l’erreur « brute ».
Partie 3: Dérivées partielles comme restrictions¶
On écrit souvent comme si c’était une quantité primitive. En réalité, cette dérivée partielle est une restriction de la dérivée totale . Deux opérations la définissent:
Injection , qui envoie un scalaire sur le vecteur (perturbation le long de l’axe uniquement)
Projection , qui extrait la composante :
La dérivée partielle est alors la composition
Autrement dit: on injecte un vecteur unitaire dans la dérivée totale, puis on projette sur la composante . Les entrées de la jacobienne ne sont pas la définition — elles en sont une conséquence.
En code, est simplement un vecteur canonique, et est l’indexation [i].
def iota(j, n):
"""Injection: iota_j(1) = e_j, le j-ième vecteur canonique de R^n."""
return jnp.eye(n)[j]
def pi(i, y):
"""Projection: pi_i(y) = y_i, la i-ième composante."""
return y[i]
# Démonstration: calculons df_1/dx_0 (a) via la restriction
e_0 = iota(0, 2) # e_0 = [1, 0]
Df_a_e0 = jax.jvp(f, (a,), (e_0,))[1] # Df(a)(e_0): colonne 0 de J_f
partial_f1_x0 = pi(1, Df_a_e0) # pi_1(Df(a)(e_0))
print(f"pi_1 . Df(a) . iota_0 = {partial_f1_x0}")
print(f"J_f(a)[1, 0] = {J_f[1, 0]}")
print(f"Accord: {jnp.isclose(partial_f1_x0, J_f[1, 0])}")pi_1 . Df(a) . iota_0 = 2.0
J_f(a)[1, 0] = 2.0
Accord: True
Exercice 2: Reconstruire la jacobienne colonne par colonne (JVP) ★¶
Un JVP donne la -ième colonne de : c’est le vecteur des réponses de toutes les sorties à une perturbation unitaire de l’entrée .
Reconstruisez la jacobienne colonne par colonne en utilisant appels à jax.jvp, un par vecteur canonique .
n = 2 # dimension d'entrée
m = 3 # dimension de sortie
J_par_jvp = jnp.zeros((m, n))
for j in range(n):
e_j = iota(j, n)
# ============================================
# TODO: Calculez Df(a)(e_j) via jax.jvp, puis placez le résultat
# dans la colonne j de J_par_jvp.
#
# col_j = jax.jvp(f, (a,), (e_j,))[1]
# J_par_jvp = J_par_jvp.at[:, j].set(col_j)
# ============================================
pass
if jnp.any(J_par_jvp != 0):
print("J_f(a) reconstruite par JVP:")
print(J_par_jvp)
print(f"\nAccord avec jax.jacobian: {jnp.allclose(J_par_jvp, J_f)}")
else:
print("Complétez la boucle!")Complétez la boucle!
Solution Exercice 2 (cliquez pour afficher)
J_par_jvp = jnp.zeros((m, n))
for j in range(n):
e_j = iota(j, n)
col_j = jax.jvp(f, (a,), (e_j,))[1]
J_par_jvp = J_par_jvp.at[:, j].set(col_j)
print("J_f(a) reconstruite par JVP:")
print(J_par_jvp)
print(f"\nAccord avec jax.jacobian: {jnp.allclose(J_par_jvp, J_f)}")Exercice 3: Reconstruire la jacobienne ligne par ligne (VJP) ★★¶
Le VJP offre une vue duale. Rappelons que jax.vjp(f, a) renvoie (f(a), vjp_fn) où vjp_fn(u) calcule .
Si on choisit (vecteur canonique de ), on obtient
Reconstruisez la jacobienne ligne par ligne en utilisant appels VJP.
Observation clé: l’exercice 2 utilise JVP (un par entrée), celui-ci utilise VJP (un par sortie). Lequel est plus économique dépend de et .
J_par_vjp = jnp.zeros((m, n))
# Un seul appel à jax.vjp suffit pour obtenir vjp_fn (la passe avant est partagée)
_, vjp_fn = jax.vjp(f, a)
for i in range(m):
e_i = iota(i, m)
# ============================================
# TODO: Calculez J_f(a)^T e_i via vjp_fn(e_i).
# Le résultat est la i-ième LIGNE de J_f(a).
#
# row_i = vjp_fn(e_i)[0]
# J_par_vjp = J_par_vjp.at[i, :].set(row_i)
# ============================================
pass
if jnp.any(J_par_vjp != 0):
print("J_f(a) reconstruite par VJP:")
print(J_par_vjp)
print(f"\nAccord avec jax.jacobian: {jnp.allclose(J_par_vjp, J_f)}")
else:
print("Complétez la boucle!")Complétez la boucle!
Solution Exercice 3 (cliquez pour afficher)
J_par_vjp = jnp.zeros((m, n))
_, vjp_fn = jax.vjp(f, a)
for i in range(m):
e_i = iota(i, m)
row_i = vjp_fn(e_i)[0]
J_par_vjp = J_par_vjp.at[i, :].set(row_i)
print("J_f(a) reconstruite par VJP:")
print(J_par_vjp)
print(f"\nAccord avec jax.jacobian: {jnp.allclose(J_par_vjp, J_f)}")Résumé: JVP avec \mathbf{e}_j extrait la colonne j, VJP avec \mathbf{e}_i extrait la ligne i. Pour f: \mathbb{R}^n \to \mathbb{R}^m, reconstruire la jacobienne complète coûte n JVP ou m VJP. Quand la sortie est scalaire (m=1), un seul VJP suffit — c’est pourquoi le mode arrière domine en apprentissage machine.
Partie 4: Règle de la chaîne — composition d’applications linéaires¶
Dans la notation de Spivak, la règle de la chaîne s’écrit
C’est une composition d’applications linéaires, pas un « produit de dérivées ». En coordonnées, cela donne le produit de matrices jacobiennes:
Définissons une seconde fonction pour composer avec :
def g(y):
"""g: R^3 -> R^2"""
return jnp.array([y[0] * y[2] + y[1]**2, jnp.exp(y[0]) - y[2]])
def h(x):
"""h = g o f: R^2 -> R^2"""
return g(f(x))
print("f(a) =", f(a))
print("h(a) = g(f(a)) =", h(a))f(a) = [3. 2. 0.84147098]
h(a) = g(f(a)) = [ 6.52441295 19.24406594]
W0409 22:49:22.814510 39649437 cpp_gen_intrinsics.cc:74] Empty bitcode string provided for eigen. Optimizations relying on this IR will be disabled.
Exercice 4: Trois façons de calculer la jacobienne d’une composition ★★¶
Calculez de trois manières et vérifiez qu’elles coïncident:
Directement:
jax.jacobian(h)(a)Produit de jacobiennes:
Composition de JVP: pour chaque , calculer par deux appels imbriqués à
jax.jvp
# Méthode 1: directe
J_h_direct = jax.jacobian(h)(a)
print("Méthode 1 (directe):")
print(J_h_direct)
# ============================================
# TODO — Méthode 2: produit de jacobiennes
# J_g_at_fa = jax.jacobian(g)(f(a)) # J_g évaluée en f(a)
# J_f_at_a = jax.jacobian(f)(a) # J_f évaluée en a
# J_h_produit = J_g_at_fa @ J_f_at_a
# ============================================
J_h_produit = None # <- Complétez
# ============================================
# TODO — Méthode 3: composition de JVP
# J_h_jvp = jnp.zeros((2, 2))
# for j in range(2):
# e_j = iota(j, 2)
# _, Df_a_ej = jax.jvp(f, (a,), (e_j,)) # Df(a)(e_j)
# _, Dg_fa_Dfej = jax.jvp(g, (f(a),), (Df_a_ej,)) # Dg(f(a))(Df(a)(e_j))
# J_h_jvp = J_h_jvp.at[:, j].set(Dg_fa_Dfej)
# ============================================
J_h_jvp = None # <- Complétez
if J_h_produit is not None and J_h_jvp is not None:
print("\nMéthode 2 (produit):")
print(J_h_produit)
print("\nMéthode 3 (JVP composés):")
print(J_h_jvp)
print(f"\n1 == 2: {jnp.allclose(J_h_direct, J_h_produit)}")
print(f"1 == 3: {jnp.allclose(J_h_direct, J_h_jvp)}")
else:
print("\nComplétez les méthodes 2 et 3!")Méthode 1 (directe):
[[11.30384889 4.84147098]
[39.63077154 20.08553692]]
Complétez les méthodes 2 et 3!
Solution Exercice 4 (cliquez pour afficher)
# Méthode 2: produit de jacobiennes
J_g_at_fa = jax.jacobian(g)(f(a))
J_f_at_a = jax.jacobian(f)(a)
J_h_produit = J_g_at_fa @ J_f_at_a
# Méthode 3: composition de JVP
J_h_jvp = jnp.zeros((2, 2))
for j in range(2):
e_j = iota(j, 2)
_, Df_a_ej = jax.jvp(f, (a,), (e_j,))
_, Dg_fa_Dfej = jax.jvp(g, (f(a),), (Df_a_ej,))
J_h_jvp = J_h_jvp.at[:, j].set(Dg_fa_Dfej)La méthode 3 illustre concrètement D(g \circ f)(\mathbf{a})(\mathbf{e}_j) = Dg(f(\mathbf{a}))(Df(\mathbf{a})(\mathbf{e}_j)) — on applique d’abord Df(\mathbf{a}), puis Dg(f(\mathbf{a})), exactement comme une composition de fonctions.
Partie 5: Le produit vecteur-jacobienne (VJP)¶
La dérivée totale est une application linéaire représentée par la matrice jacobienne . Le JVP multiplie cette matrice par un vecteur à droite: .
Le VJP (vector-Jacobian product) effectue l’opération duale: on prend un vecteur et on le multiplie à gauche de la jacobienne:
Le résultat est un vecteur de , de la même taille que l’entrée de . Le vecteur est appelé cotangent. Il représente un signal qui arrive de la sortie, et le VJP le traduit vers l’entrée. En rétropropagation, chaque couche reçoit un cotangent et le propage en arrière grâce à son VJP.
| Direction | Opération | Rôle |
|---|---|---|
| JVP | Propage un signal vers l’avant (entrée → sortie) | |
| VJP | Propage un signal vers l’arrière (sortie → entrée) |
En JAX, jax.vjp(f, a) renvoie (f(a), vjp_fn), et vjp_fn(u) calcule .
Dans certains manuels, le VJP est écrit (transposer la matrice puis multiplier à droite) plutôt que (multiplier à gauche). Les deux donnent le même vecteur. La seconde écriture évite de transposer et correspond littéralement au nom vector-Jacobian product. En algèbre linéaire, cette opération est l’adjoint de , noté .
# Démonstration: VJP de f au point a
f_a, vjp_fn = jax.vjp(f, a)
u = jnp.array([1.0, 0.0, -0.5])
# VJP via JAX
vjp_jax = vjp_fn(u)[0]
# VJP via produit vecteur-jacobienne explicite
vjp_explicit = u @ J_f
print("VJP (JAX): ", vjp_jax)
print("u^T @ J_f: ", vjp_explicit)
print("Accord:", jnp.allclose(vjp_jax, vjp_explicit))VJP (JAX): [1.72984885 1. ]
u^T @ J_f: [1.72984885 1. ]
Accord: True
Méthode systématique pour dériver une règle VJP¶
Le VJP d’une fonction au point est le produit . On pourrait calculer en entier puis effectuer le produit matriciel, mais la jacobienne est souvent trop grande pour être stockée. On cherche donc à écrire sans former la matrice. Pour beaucoup d’opérations courantes, le résultat se réduit à quelques multiplications élément par élément.
Recette en 3 étapes pour toute opération :
Calculer les dérivées partielles et les organiser dans la matrice .
Calculer . La -ième composante du résultat est .
Simplifier: chercher une formule directe (produit élément par élément, produit matrice-vecteur, etc.) qui donne le même résultat sans construire la matrice.
Deux exemples avec des valeurs numériques illustrent cette recette.
Exemple A: exponentielle élément par élément¶
Soit , c’est-à-dire . Prenons et .
Étape 1. Chaque sortie ne dépend que de , donc les dérivées hors-diagonale sont nulles:
Étape 2. On multiplie le vecteur ligne par chaque colonne de . Puisque est diagonale, chaque colonne n’a qu’une entrée non nulle:
Étape 3. Chaque composante du résultat est , ce qui donne une multiplication élément par élément:
Le coût est au lieu de si on avait formé la matrice entière.
Exemple B: somme (réduction)¶
Soit , qui prend un vecteur de nombres et renvoie un scalaire. Prenons et . Ici est un scalaire puisque la sortie l’est aussi.
Étape 1. La sortie est un scalaire, donc la jacobienne est un vecteur ligne:
Étape 2. Un scalaire fois un vecteur ligne:
Étape 3. Le scalaire est copié sur chaque composante:
La somme contracte valeurs en une seule; son VJP fait l’inverse et diffuse le scalaire en copies. Ce patron réduction/diffusion revient constamment en rétropropagation.
Les exercices suivants appliquent cette même recette à trois opérations plus riches.
Exercice 5: Dériver et vérifier des règles VJP ★★¶
Appliquez la recette en 4 étapes pour dériver la règle VJP de chaque opération ci-dessous, puis vérifiez numériquement en implémentant la formule manuelle et en comparant avec jax.vjp.
(a) Carré élément par élément:
La jacobienne est . La règle VJP:
def phi(x):
return x ** 2
a_test = jnp.array([1.0, 2.0, 3.0])
u_test = jnp.array([0.1, -0.2, 0.5])
# ============================================
# TODO (a): Implémentez la règle VJP manuellement
# vjp_manual_a = 2 * a_test * u_test
# ============================================
vjp_manual_a = None # <- Complétez
# Vérification avec JAX
_, vjp_phi = jax.vjp(phi, a_test)
vjp_jax_a = vjp_phi(u_test)[0]
if vjp_manual_a is not None:
print("(a) Carré élément par élément:")
print(f" VJP manuelle: {vjp_manual_a}")
print(f" VJP JAX: {vjp_jax_a}")
print(f" Accord: {jnp.allclose(vjp_manual_a, vjp_jax_a)}")
else:
print("Complétez la VJP manuelle (a)!")Complétez la VJP manuelle (a)!
(b) Produit matrice-vecteur: avec fixe.
La jacobienne par rapport à est (application linéaire constante). La règle VJP:
W = jnp.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) # 3x2
z = jnp.array([0.5, -1.0])
u_b = jnp.array([1.0, 0.0, -0.5])
def matmul_z(z):
return W @ z
# ============================================
# TODO (b): VJP manuelle du produit matrice-vecteur
# vjp_manual_b = W.T @ u_b
# ============================================
vjp_manual_b = None # <- Complétez
# Vérification avec JAX
_, vjp_matmul = jax.vjp(matmul_z, z)
vjp_jax_b = vjp_matmul(u_b)[0]
if vjp_manual_b is not None:
print("(b) Produit matrice-vecteur:")
print(f" VJP manuelle: {vjp_manual_b}")
print(f" VJP JAX: {vjp_jax_b}")
print(f" Accord: {jnp.allclose(vjp_manual_b, vjp_jax_b)}")
else:
print("Complétez la VJP manuelle (b)!")Complétez la VJP manuelle (b)!
(c) ★★★ Softmax: , .
La jacobienne a pour entrées , ce qui donne la règle VJP:
Ce concept est subtil. Prenez le temps de dérouler le produit à la main pour voir comment cette formule émerge.
def softmax(a):
a_stable = a - jnp.max(a)
exp_a = jnp.exp(a_stable)
return exp_a / jnp.sum(exp_a)
a_sm = jnp.array([2.0, 1.0, 0.1])
u_sm = jnp.array([1.0, -0.5, 0.3])
# ============================================
# TODO (c): VJP manuelle du softmax
# s = softmax(a_sm)
# vjp_manual_c = s * (u_sm - jnp.dot(s, u_sm))
# ============================================
vjp_manual_c = None # <- Complétez
# Vérification avec JAX
_, vjp_softmax = jax.vjp(softmax, a_sm)
vjp_jax_c = vjp_softmax(u_sm)[0]
if vjp_manual_c is not None:
print("(c) Softmax:")
print(f" VJP manuelle: {vjp_manual_c}")
print(f" VJP JAX: {vjp_jax_c}")
print(f" Accord: {jnp.allclose(vjp_manual_c, vjp_jax_c)}")
else:
print("Complétez la VJP manuelle (c)!")Complétez la VJP manuelle (c)!
Solution Exercice 5 (cliquez pour afficher)
# (a)
vjp_manual_a = 2 * a_test * u_test
# (b)
vjp_manual_b = u_b @ W
# (c)
s = softmax(a_sm)
vjp_manual_c = s * (u_sm - jnp.dot(s, u_sm))(a) Carré élément par élément, \varphi(\mathbf{x}) = \mathbf{x}^{\odot 2}:
\frac{\partial (x_i^2)}{\partial x_j} = 2a_i \delta_{ij}, donc \mathbf{J} = \operatorname{diag}(2a_1, \ldots, 2a_n).
(\mathbf{u}^\top \mathbf{J})_j = u_j \cdot 2a_j (matrice diagonale: chaque colonne n’a qu’une entrée).
\text{VJP} = 2\mathbf{a} \odot \mathbf{u}, un produit de Hadamard en O(n) au lieu de O(n^2).
(b) Produit matrice-vecteur, f(\mathbf{z}) = W\mathbf{z}:
La fonction est linéaire, donc \mathbf{J} = W \in \mathbb{R}^{m \times n}.
\mathbf{u}^\top W, un produit vecteur-matrice.
\text{VJP} = \mathbf{u}^\top W en O(mn).
En code, u_b @ W donne le même résultat que W.T @ u_b.
(c) Softmax, \sigma_i(\mathbf{a}) = \frac{e^{a_i}}{\sum_j e^{a_j}}:
\frac{\partial \sigma_i}{\partial a_j} = \sigma_i(\delta_{ij} - \sigma_j). En forme matricielle: \mathbf{J} = \operatorname{diag}(\boldsymbol{\sigma}) - \boldsymbol{\sigma}\boldsymbol{\sigma}^\top.
(\mathbf{u}^\top \mathbf{J})_j = \sum_i u_i \sigma_i(\delta_{ij} - \sigma_j) = u_j \sigma_j - \sigma_j \sum_i u_i \sigma_i = \sigma_j(u_j - \langle \boldsymbol{\sigma}, \mathbf{u} \rangle).
\text{VJP} = \boldsymbol{\sigma} \odot (\mathbf{u} - \langle \boldsymbol{\sigma}, \mathbf{u}\rangle \, \mathbf{1}), soit O(K) au lieu de O(K^2).
La jacobienne du softmax est symétrique (\mathbf{J} = \mathbf{J}^\top), donc \mathbf{u}^\top \mathbf{J} = (\mathbf{J}\mathbf{u})^\top. Ce n’est pas le cas en général.
Exercice 6: Gradients par composition de VJP ★★¶
L’exercice 5 a établi les règles VJP pour des opérations individuelles. Ici, on pratique la compétence complémentaire: décomposer une fonction scalaire en chaîne de primitives, puis composer les VJP en ordre inverse pour obtenir le gradient.
La méthode est toujours la même. Soit :
Passe avant: identifier les valeurs intermédiaires , , ..., .
Passe arrière: partir de (cotangent initial pour une perte scalaire) et propager: .
Le gradient est .
Dérivez le gradient sur papier en composant les règles VJP de l’exercice 5, puis vérifiez avec jax.grad.
(a) ★ Norme au carré:
Décomposition:
(b) ★★ Log-sum-exp:
Décomposition:
(c) ★★ Moindres carrés: avec et fixes.
Décomposition:
(d) ★★★ Entropie croisée: avec la classe cible et le softmax.
Décomposition:
x_drill = jnp.array([1.0, 2.0, -0.5, 0.3])
# --- (a) Norme au carré ---
def loss_a(x):
return jnp.sum(x ** 2)
# ============================================
# TODO (a): Écrivez le gradient obtenu par composition de VJP
# grad_manual_a = ???
# ============================================
grad_manual_a = None # <- Complétez
grad_jax_a = jax.grad(loss_a)(x_drill)
if grad_manual_a is not None:
print(f"(a) Manuel: {grad_manual_a}")
print(f" JAX: {grad_jax_a}")
print(f" Accord: {jnp.allclose(grad_manual_a, grad_jax_a)}\n")
else:
print("(a) Complétez grad_manual_a!\n")
# --- (b) Log-sum-exp ---
def loss_b(x):
return jnp.log(jnp.sum(jnp.exp(x)))
# ============================================
# TODO (b): Écrivez le gradient obtenu par composition de VJP
# grad_manual_b = ???
# ============================================
grad_manual_b = None # <- Complétez
grad_jax_b = jax.grad(loss_b)(x_drill)
if grad_manual_b is not None:
print(f"(b) Manuel: {grad_manual_b}")
print(f" JAX: {grad_jax_b}")
print(f" Accord: {jnp.allclose(grad_manual_b, grad_jax_b)}\n")
else:
print("(b) Complétez grad_manual_b!\n")
# --- (c) Moindres carrés ---
key_drill = jax.random.PRNGKey(7)
A_drill = jax.random.normal(key_drill, (3, 4))
b_drill = jnp.array([1.0, -1.0, 0.5])
def loss_c(x):
r = A_drill @ x - b_drill
return jnp.sum(r ** 2)
# ============================================
# TODO (c): Écrivez le gradient obtenu par composition de VJP
# grad_manual_c = ???
# ============================================
grad_manual_c = None # <- Complétez
grad_jax_c = jax.grad(loss_c)(x_drill)
if grad_manual_c is not None:
print(f"(c) Manuel: {grad_manual_c}")
print(f" JAX: {grad_jax_c}")
print(f" Accord: {jnp.allclose(grad_manual_c, grad_jax_c)}\n")
else:
print("(c) Complétez grad_manual_c!\n")
# --- (d) Entropie croisée ---
k_true = 1 # classe cible
def loss_d(z):
return -jnp.log(softmax(z)[k_true])
# ============================================
# TODO (d): Écrivez le gradient obtenu par composition de VJP
# grad_manual_d = ???
# ============================================
grad_manual_d = None # <- Complétez
grad_jax_d = jax.grad(loss_d)(x_drill)
if grad_manual_d is not None:
print(f"(d) Manuel: {grad_manual_d}")
print(f" JAX: {grad_jax_d}")
print(f" Accord: {jnp.allclose(grad_manual_d, grad_jax_d, atol=1e-7)}")
else:
print("(d) Complétez grad_manual_d!")(a) Complétez grad_manual_a!
(b) Complétez grad_manual_b!
(c) Complétez grad_manual_c!
(d) Complétez grad_manual_d!
Solution Exercice 6 (cliquez pour afficher)
# (a)
grad_manual_a = 2 * x_drill
# (b)
grad_manual_b = softmax(x_drill)
# (c)
r = A_drill @ x_drill - b_drill
grad_manual_c = 2 * A_drill.T @ r
# (d)
s = softmax(x_drill)
e_k = jnp.zeros_like(x_drill).at[k_true].set(1.0)
grad_manual_d = s - e_kDérivation (a) — \ell(\mathbf{x}) = \sum_i x_i^2:
Passe arrière (\bar{\ell} = 1):
[D\text{sum}]^*(1) = 1 \cdot \mathbf{1} = \mathbf{1} (diffusion du scalaire)
[D(\cdot)^2(\mathbf{a})]^*(\mathbf{1}) = 2\mathbf{a} \odot \mathbf{1} = 2\mathbf{a}
Résultat: \nabla \ell(\mathbf{a}) = 2\mathbf{a}.
Dérivation (b) — \ell(\mathbf{x}) = \log(\sum_i e^{x_i}):
Passe arrière (\bar{\ell} = 1):
[D\log(S)]^*(1) = 1/S
[D\text{sum}]^*(1/S) = (1/S) \cdot \mathbf{1}
[D\exp(\mathbf{a})]^*\bigl(\frac{1}{S}\mathbf{1}\bigr) = e^{\mathbf{a}} \odot \frac{1}{S}\mathbf{1} = \frac{e^{\mathbf{a}}}{\sum_j e^{a_j}} = \boldsymbol{\sigma}(\mathbf{a})
Résultat: \nabla \text{lse}(\mathbf{a}) = \text{softmax}(\mathbf{a}). Le gradient du log-sum-exp est le softmax.
Dérivation (c) — \ell(\mathbf{x}) = \|A\mathbf{x} - \mathbf{b}\|^2:
Passe arrière (\bar{\ell} = 1):
[D\text{sum}]^*(1) = \mathbf{1}
[D(\cdot)^2(\mathbf{r})]^*(\mathbf{1}) = 2\mathbf{r}
[D(-\mathbf{b})]^*(2\mathbf{r}) = 2\mathbf{r} (translation: VJP = identité)
[D(A \cdot)]^*(2\mathbf{r}) = A^\top (2\mathbf{r}) = 2A^\top(A\mathbf{a} - \mathbf{b})
Résultat: \nabla \ell(\mathbf{a}) = 2A^\top(A\mathbf{a} - \mathbf{b}). On reconnaît l’équation normale de la régression linéaire.
Dérivation (d) — \ell(\mathbf{z}) = -\log(\sigma(\mathbf{z})_k):
Passe arrière (\bar{\ell} = 1):
[D(-\log)(\sigma_k)]^*(1) = -1/\sigma_k
[D\pi_k]^*(-1/\sigma_k) = (-1/\sigma_k) \, \mathbf{e}_k. (L’adjoint de la projection est l’injection: \pi_k^*(c) = c \, \mathbf{e}_k.)
[D\sigma(\mathbf{a})]^*\bigl((-1/\sigma_k) \mathbf{e}_k\bigr): on applique la règle du softmax \boldsymbol{\sigma} \odot (\mathbf{u} - \langle \boldsymbol{\sigma}, \mathbf{u}\rangle \mathbf{1}) avec \mathbf{u} = (-1/\sigma_k)\mathbf{e}_k.
\langle \boldsymbol{\sigma}, \mathbf{u} \rangle = \sigma_k \cdot (-1/\sigma_k) = -1
\mathbf{u} - (-1)\mathbf{1} = (-1/\sigma_k)\mathbf{e}_k + \mathbf{1}
\boldsymbol{\sigma} \odot \bigl((-1/\sigma_k)\mathbf{e}_k + \mathbf{1}\bigr) = \boldsymbol{\sigma} - \mathbf{e}_k
Résultat: \nabla \ell(\mathbf{a}) = \boldsymbol{\sigma}(\mathbf{a}) - \mathbf{e}_k. C’est le gradient classique de l’entropie croisée avec softmax: la prédiction moins la cible.
Partie 6: Composition de VJP et coût computationnel¶
En prenant l’adjoint de la règle de la chaîne , et en utilisant , on obtient
L’ordre s’inverse: on applique d’abord l’adjoint de la fonction extérieure , puis celui de la fonction intérieure . C’est la rétropropagation: le signal se propage de la sortie vers l’entrée.
Pourquoi le nombre de passes dépend des dimensions¶
Rappelons les exercices 2 et 3: un JVP avec donne la colonne de la jacobienne, un VJP avec donne la ligne . Pour reconstruire la jacobienne complète , il faut donc:
passes JVP (une par colonne), ou
passes VJP (une par ligne)
Voyons ce que cela donne concrètement sur un petit exemple.
# --- Démonstration: JVP et VJP avec vecteurs canoniques, pas à pas ---
# Reprenons f: R^2 -> R^3 (n=2 entrées, m=3 sorties)
print("f: R^2 -> R^3 (n=2 entrées, m=3 sorties)")
print("=" * 55)
# La jacobienne complète (pour référence)
J_ref = jax.jacobian(f)(a)
print(f"\nJacobienne complète (référence):\n{J_ref}\n")
# --- Mode avant (JVP): n=2 passes pour remplir 2 colonnes ---
print("Mode avant: 2 JVP (un par vecteur canonique de R^n)")
print("-" * 55)
for j in range(2):
e_j = iota(j, 2)
_, col_j = jax.jvp(f, (a,), (e_j,))
print(f" JVP(f, a, e_{j}) = Df(a)(e_{j}) = {col_j} <- colonne {j} de J")
# --- Mode arrière (VJP): m=3 passes pour remplir 3 lignes ---
print(f"\nMode arrière: 3 VJP (un par vecteur canonique de R^m)")
print("-" * 55)
_, vjp_f_fn = jax.vjp(f, a)
for i in range(3):
e_i = iota(i, 3)
row_i = vjp_f_fn(e_i)[0]
print(f" VJP(f, a, e_{i}) = e_{i}^T J = {row_i} <- ligne {i} de J")
print(f"\nBilan: {2} JVP ou {3} VJP pour la même jacobienne 3×2.")
print("Ici n < m, donc le mode avant (JVP) requiert moins de passes.")f: R^2 -> R^3 (n=2 entrées, m=3 sorties)
=======================================================
Jacobienne complète (référence):
[[2. 1. ]
[2. 1. ]
[0.54030231 0. ]]
Mode avant: 2 JVP (un par vecteur canonique de R^n)
-------------------------------------------------------
JVP(f, a, e_0) = Df(a)(e_0) = [2. 2. 0.54030231] <- colonne 0 de J
JVP(f, a, e_1) = Df(a)(e_1) = [1. 1. 0.] <- colonne 1 de J
Mode arrière: 3 VJP (un par vecteur canonique de R^m)
-------------------------------------------------------
VJP(f, a, e_0) = e_0^T J = [2. 1.] <- ligne 0 de J
VJP(f, a, e_1) = e_1^T J = [2. 1.] <- ligne 1 de J
VJP(f, a, e_2) = e_2^T J = [0.54030231 0. ] <- ligne 2 de J
Bilan: 2 JVP ou 3 VJP pour la même jacobienne 3×2.
Ici n < m, donc le mode avant (JVP) requiert moins de passes.
# --- Cas inverse: beaucoup d'entrées, sortie scalaire (comme une perte) ---
def loss(x):
"""loss: R^5 -> R^1 (perte scalaire)"""
return jnp.sum(jnp.sin(x) * x)
x_demo = jnp.array([1.0, 2.0, 3.0, 4.0, 5.0])
n_demo, m_demo = 5, 1
print(f"\nloss: R^{n_demo} -> R^{m_demo} (n={n_demo} entrées, m={m_demo} sortie)")
print("=" * 55)
# Mode avant: n=5 passes (une par entrée)
print(f"\nMode avant: {n_demo} JVP nécessaires")
print("-" * 55)
grad_jvp = jnp.zeros(n_demo)
for j in range(n_demo):
e_j = iota(j, n_demo)
_, djvp = jax.jvp(loss, (x_demo,), (e_j,))
grad_jvp = grad_jvp.at[j].set(djvp)
print(f" JVP(loss, x, e_{j}) = {djvp:.4f} <- composante {j} du gradient")
# Mode arrière: m=1 seule passe!
print(f"\nMode arrière: {m_demo} seul VJP suffit")
print("-" * 55)
grad_vjp = jax.grad(loss)(x_demo)
print(f" VJP(loss, x, 1.0) = {grad_vjp} <- gradient complet en une passe!")
print(f"\nAccord: {jnp.allclose(grad_jvp, grad_vjp)}")
print(f"\nBilan: {n_demo} JVP vs {m_demo} VJP. Ratio = {n_demo}/{m_demo} = {n_demo}×")
print(f"Pour un réseau avec n = 10^6 paramètres et perte scalaire:")
print(f" Mode avant: 10^6 passes JVP")
print(f" Mode arrière: 1 passe VJP <- c'est la rétropropagation!")
loss: R^5 -> R^1 (n=5 entrées, m=1 sortie)
=======================================================
Mode avant: 5 JVP nécessaires
-------------------------------------------------------
JVP(loss, x, e_0) = 1.3818 <- composante 0 du gradient
JVP(loss, x, e_1) = 0.0770 <- composante 1 du gradient
JVP(loss, x, e_2) = -2.8289 <- composante 2 du gradient
JVP(loss, x, e_3) = -3.3714 <- composante 3 du gradient
JVP(loss, x, e_4) = 0.4594 <- composante 4 du gradient
Mode arrière: 1 seul VJP suffit
-------------------------------------------------------
VJP(loss, x, 1.0) = [ 1.38177329 0.07700375 -2.82885748 -3.37137698 0.45938665] <- gradient complet en une passe!
Accord: True
Bilan: 5 JVP vs 1 VJP. Ratio = 5/1 = 5×
Pour un réseau avec n = 10^6 paramètres et perte scalaire:
Mode avant: 10^6 passes JVP
Mode arrière: 1 passe VJP <- c'est la rétropropagation!
Le tableau suivant résume le nombre de passes nécessaires pour différentes configurations:
| Fonction | (entrées) | (sorties) | Passes JVP | Passes VJP | Mode le plus économique |
|---|---|---|---|---|---|
| 2 | 3 | 2 | 3 | JVP (mode avant) | |
| Perte scalaire: | 5 | 1 | 5 | 1 | VJP (mode arrière) |
| Réseau de neurones: | 106 | 1 | 106 | 1 | VJP (mode arrière) |
| Couche cachée: | 1 | 100 | 1 | 100 | JVP (mode avant) |
Règle générale: le mode arrière (VJP) domine quand . En apprentissage machine, la perte est toujours scalaire () et le nombre de paramètres est grand (). C’est pourquoi la rétropropagation — qui est une composition de VJP — est l’algorithme standard.
Exercice 7: Composer des VJP ★★¶
En reprenant et le cotangent (dimension de sortie de ), calculez de trois façons.
u_h = jnp.array([1.0, -1.0])
# Méthode 1: VJP directe de h
_, vjp_h = jax.vjp(h, a)
result_1 = vjp_h(u_h)[0]
print("Méthode 1 (VJP directe):", result_1)
# ============================================
# TODO — Méthode 2: produit matriciel explicite u^T J_g J_f
# J_g_fa = jax.jacobian(g)(f(a))
# J_f_a = jax.jacobian(f)(a)
# result_2 = u_h @ J_g_fa @ J_f_a
# ============================================
result_2 = None # <- Complétez
# ============================================
# TODO — Méthode 3: composition d'appels VJP
# Passe avant
# f_a = f(a)
# _, vjp_g_fn = jax.vjp(g, f_a) # VJP de g en f(a)
# _, vjp_f_fn = jax.vjp(f, a) # VJP de f en a
#
# Passe arrière (ordre inversé!)
# u_mid = vjp_g_fn(u_h)[0] # u^T J_g
# result_3 = vjp_f_fn(u_mid)[0] # (u^T J_g) J_f
# ============================================
result_3 = None # <- Complétez
if result_2 is not None and result_3 is not None:
print("Méthode 2 (u^T J produit):", result_2)
print("Méthode 3 (VJP composés):", result_3)
print(f"\n1 == 2: {jnp.allclose(result_1, result_2)}")
print(f"1 == 3: {jnp.allclose(result_1, result_3)}")
else:
print("\nComplétez les méthodes 2 et 3!")Méthode 1 (VJP directe): [-28.32692265 -15.24406594]
Complétez les méthodes 2 et 3!
Solution Exercice 7 (cliquez pour afficher)
# Méthode 2
J_g_fa = jax.jacobian(g)(f(a))
J_f_a = jax.jacobian(f)(a)
result_2 = J_f_a.T @ (J_g_fa.T @ u_h)
# Méthode 3
f_a = f(a)
_, vjp_g_fn = jax.vjp(g, f_a)
_, vjp_f_fn = jax.vjp(f, a)
u_mid = vjp_g_fn(u_h)[0] # [Dg(f(a))]*(u) — signal intermédiaire
result_3 = vjp_f_fn(u_mid)[0] # [Df(a)]*(u_mid) — signal à l'entréeNotez l’ordre inversé dans la méthode 3: on propage d’abord à travers g (la dernière fonction appliquée dans la passe avant), puis à travers f.
Exercice 8: Coût computationnel — mode avant vs mode arrière ★★★¶
Considérons une chaîne de trois couches (inspirée d’un réseau de neurones avec perte scalaire):
La composition est .
Les jacobiennes sont: , , .
Le gradient complet est (un vecteur).
(a) Le coût d’un produit matrice-vecteur avec est multiplications. Comptez les multiplications pour les deux ordres d’évaluation:
| Ordre | Opérations | Multiplications |
|---|---|---|
| Mode avant (JVP): | , puis , puis | ? + ? + ? = ? par tangent |
| Mode arrière (VJP): | , puis , puis | ? + ? + ? = ? par cotangent |
Pour le gradient complet (): le mode avant requiert passes JVP (une par ), le mode arrière requiert passe VJP. Quel est le rapport des coûts totaux?
# ============================================
# TODO (a): Complétez le décompte des multiplications
# ============================================
# Mode avant (JVP) — une passe:
# J1 @ v: 50 x 100 = ?
# J2 @ (...): 20 x 50 = ?
# J3 @ (...): 1 x 20 = ?
cout_jvp_une_passe = None # <- Complétez (somme)
cout_jvp_total = None # <- Complétez (x 100 passes)
# Mode arrière (VJP) — une passe:
# J3^T @ u: 20 x 1 = ?
# J2^T @ (...): 50 x 20 = ?
# J1^T @ (...): 100 x 50 = ?
cout_vjp_une_passe = None # <- Complétez (somme)
cout_vjp_total = None # <- Complétez (x 1 passe)
if cout_jvp_total is not None:
print(f"Mode avant: {cout_jvp_une_passe} mult/passe × 100 passes = {cout_jvp_total}")
print(f"Mode arrière: {cout_vjp_une_passe} mult/passe × 1 passe = {cout_vjp_total}")
print(f"Rapport: {cout_jvp_total / cout_vjp_total:.0f}× — égal à n (dimension d'entrée)")
else:
print("Complétez le décompte!")Complétez le décompte!
(b) Vérifions avec JAX. On crée une chaîne de couches linéaires (avec tanh) et on compare le temps pour calculer le gradient de deux façons: 100 JVP (mode avant) vs un seul jax.grad (mode arrière).
import time
key = jax.random.PRNGKey(42)
keys = jax.random.split(key, 3)
W1 = jax.random.normal(keys[0], (50, 100)) * 0.1
W2 = jax.random.normal(keys[1], (20, 50)) * 0.1
W3 = jax.random.normal(keys[2], (1, 20)) * 0.1
def chain(x):
z1 = jnp.tanh(W1 @ x)
z2 = jnp.tanh(W2 @ z1)
return (W3 @ z2)[0] # scalaire
x0 = jax.random.normal(jax.random.PRNGKey(0), (100,))
# Échauffement (compilation JIT)
_ = jax.grad(chain)(x0)
_ = jax.jvp(chain, (x0,), (jnp.ones(100),))
# Mode arrière: un seul jax.grad
t0 = time.perf_counter()
for _ in range(100):
grad_reverse = jax.grad(chain)(x0)
t_reverse = (time.perf_counter() - t0) / 100
# Mode avant: 100 JVP pour reconstruire le gradient
t0 = time.perf_counter()
for _ in range(100):
grad_forward = jnp.zeros(100)
for j in range(100):
e_j = iota(j, 100)
_, djvp = jax.jvp(chain, (x0,), (e_j,))
grad_forward = grad_forward.at[j].set(djvp)
t_forward = (time.perf_counter() - t0) / 100
print(f"Mode arrière (1 VJP): {t_reverse*1000:.2f} ms")
print(f"Mode avant (100 JVP): {t_forward*1000:.2f} ms")
print(f"Rapport: {t_forward/t_reverse:.1f}×")
print(f"\nGradients identiques: {jnp.allclose(grad_reverse, grad_forward, atol=1e-5)}")Mode arrière (1 VJP): 12.96 ms
Mode avant (100 JVP): 241.23 ms
Rapport: 18.6×
Gradients identiques: True
(c) Inversons les dimensions: , , . Maintenant la sortie est en dimension élevée et l’entrée est scalaire.
Question: Combien de JVP faut-il pour la jacobienne complète? Combien de VJP? Quel mode domine?
keys_inv = jax.random.split(jax.random.PRNGKey(99), 3)
W1_inv = jax.random.normal(keys_inv[0], (20, 1)) * 0.1
W2_inv = jax.random.normal(keys_inv[1], (50, 20)) * 0.1
W3_inv = jax.random.normal(keys_inv[2], (100, 50)) * 0.1
def chain_inv(x):
"""R^1 -> R^100: sortie de grande dimension."""
z1 = jnp.tanh(W1_inv @ x)
z2 = jnp.tanh(W2_inv @ z1)
return W3_inv @ z2
x0_inv = jnp.array([1.0])
# Mode avant: 1 seul JVP suffit (n=1)
_, jac_forward = jax.jvp(chain_inv, (x0_inv,), (jnp.ones(1),))
# Mode arrière: 100 VJP nécessaires (m=100)
_, vjp_inv_fn = jax.vjp(chain_inv, x0_inv)
jac_reverse = jnp.zeros((100, 1))
for i in range(100):
e_i = iota(i, 100)
row_i = vjp_inv_fn(e_i)[0]
jac_reverse = jac_reverse.at[i, :].set(row_i)
print(f"Jacobienne (mode avant, 1 JVP): forme {jac_forward.shape}")
print(f"Jacobienne (mode arrière, 100 VJP): forme {jac_reverse.shape}")
print(f"Accord: {jnp.allclose(jac_forward.reshape(-1), jac_reverse.reshape(-1), atol=1e-5)}")
print(f"\nRègle générale: n < m → mode avant gagne; n > m → mode arrière gagne.")
print(f"En apprentissage machine, la perte est scalaire (m=1): le mode arrière domine toujours.")Jacobienne (mode avant, 1 JVP): forme (100,)
Jacobienne (mode arrière, 100 VJP): forme (100, 1)
Accord: True
Règle générale: n < m → mode avant gagne; n > m → mode arrière gagne.
En apprentissage machine, la perte est scalaire (m=1): le mode arrière domine toujours.
Solution Exercice 8 (a) (cliquez pour afficher)
| Ordre | Opérations | Multiplications |
|---|---|---|
| Mode avant (JVP) | 50 \times 100 + 20 \times 50 + 1 \times 20 | 5000 + 1000 + 20 = 6020 par tangent |
| Mode arrière (VJP) | 20 \times 1 + 50 \times 20 + 100 \times 50 | 20 + 1000 + 5000 = 6020 par cotangent |
Le coût par passe est le même! La différence vient du nombre de passes:
Mode avant: 100 passes (une par \mathbf{e}_j) → 100 \times 6020 = 602\,000
Mode arrière: 1 passe (un seul cotangent u = 1) → 1 \times 6020 = 6020
Rapport: 602\,000 / 6020 = 100 = n, la dimension d’entrée
cout_jvp_une_passe = 50*100 + 20*50 + 1*20 # = 6020
cout_jvp_total = cout_jvp_une_passe * 100 # = 602000
cout_vjp_une_passe = 20*1 + 50*20 + 100*50 # = 6020
cout_vjp_total = cout_vjp_une_passe * 1 # = 6020Partie 7: Vérification par différences finies¶
Les différences finies fournissent un outil de débogage indispensable. L’approximation par différences centrées:
donne une erreur (bien meilleure que la différence avant, qui est ). Le choix de est un compromis: trop grand → erreur de troncature; trop petit → erreurs d’arrondi en virgule flottante. En pratique, en float64 est un bon choix.
Exercice 9: Vérificateur de jacobienne par différences finies ★¶
Implémentez une fonction qui approxime la jacobienne par différences finies centrées, puis comparez avec la jacobienne exacte de au point .
def finite_diff_jacobian(f, x, eps=1e-7):
"""Jacobienne par différences finies centrées."""
x = jnp.asarray(x, dtype=jnp.float64)
f_x = f(x)
n = x.shape[0]
m = f_x.shape[0]
J = np.zeros((m, n))
# ============================================
# TODO: Pour chaque j de 0 à n-1:
# e_j = iota(j, n)
# J[:, j] = (f(x + eps * e_j) - f(x - eps * e_j)) / (2 * eps)
# ============================================
return J
# Test
J_fd = finite_diff_jacobian(f, a)
if np.any(J_fd != 0):
print("Jacobienne par différences finies:")
print(J_fd)
print(f"\nJacobienne exacte (JAX):")
print(np.array(J_f))
print(f"\nErreur max: {np.max(np.abs(J_fd - np.array(J_f))):.2e}")
else:
print("Complétez finite_diff_jacobian!")Complétez finite_diff_jacobian!
Solution Exercice 9 (cliquez pour afficher)
def finite_diff_jacobian(f, x, eps=1e-7):
x = jnp.asarray(x, dtype=jnp.float64)
f_x = f(x)
n = x.shape[0]
m = f_x.shape[0]
J = np.zeros((m, n))
for j in range(n):
e_j = iota(j, n)
J[:, j] = (f(x + eps * e_j) - f(x - eps * e_j)) / (2 * eps)
return JTraçons l’erreur en fonction de pour observer le compromis troncature/arrondi.
epsilons = np.logspace(-1, -15, 30)
errors = []
J_exact = np.array(jax.jacobian(f)(a))
for eps in epsilons:
J_fd_eps = finite_diff_jacobian(f, a, eps=eps)
errors.append(np.max(np.abs(J_fd_eps - J_exact)))
plt.loglog(epsilons, errors, 'o-')
plt.xlabel('$\\varepsilon$')
plt.ylabel('Erreur max $\\|\\mathbf{J}_{\\mathrm{fd}} - \\mathbf{J}_{\\mathrm{exact}}\\|_\\infty$')
plt.title('Compromis troncature/arrondi des différences finies centrées')
plt.axvline(1e-7, color='red', linestyle='--', alpha=0.5, label='$\\varepsilon = 10^{-7}$')
plt.legend()
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()
Récapitulatif¶
Ce TP a relié la notation de Spivak au code JAX. Le tableau suivant résume les correspondances.
| Concept (Spivak) | Notation | Code JAX | Interprétation |
|---|---|---|---|
| Dérivée totale | jax.jacobian(f)(a) | Application linéaire | |
| Application à un tangent | jax.jvp(f, (a,), (v,))[1] | JVP, mode avant | |
| VJP | jax.vjp(f, a)[1](u)[0] | VJP, mode arrière | |
| Dérivée partielle | jax.jacobian(f)(a)[i, j] | Restriction de | |
| Règle de la chaîne | Composer jax.jvp ou jax.vjp | Composition d’opérateurs linéaires |
À retenir:
La dérivée totale est une application linéaire, pas une matrice. La matrice jacobienne la représente.
Les dérivées partielles sont des restrictions: .
La règle de la chaîne est une composition d’applications linéaires. Pour le VJP, l’ordre de composition s’inverse.
Le VJP est le produit vecteur-jacobienne: .
Pour (perte scalaire), un seul VJP donne le gradient complet, contre JVP. C’est le mode arrière (rétropropagation).
Pour aller plus loin: Chapitre 7: Réseaux de neurones, sections « Règles VJP » et « Implémentation minimale ».