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.

TP6: Dérivée totale, VJP et différentiation automatique

IFT3395/IFT6390 — Fondements de l’apprentissage machine

Open In Colab

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 Df(a)Df(\mathbf{a}) comme une application linéaire, pas une matrice

  • Relier les dérivées partielles aux restrictions de Df(a)Df(\mathbf{a}) (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 — ∂y∂x\frac{\partial y}{\partial x}, ∂f∂xj\frac{\partial f}{\partial x_j} — 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 ∂f∂x1\frac{\partial f}{\partial x_1}, 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: Df(a)Df(\mathbf{a}) désigne la dérivée de la fonction ff, évaluée au point a\mathbf{a}. 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:

  1. Elle correspond à JAX. L’appel jax.jvp(f, (a,), (v,)) applique littéralement l’application linéaire Df(a)Df(\mathbf{a}) au vecteur v\mathbf{v}. La fonction f est le premier argument, le point a\mathbf{a} le deuxième, la direction v\mathbf{v} le troisième. Pas de « ∂\partial sortie / ∂\partial entrée » — juste: fonction, point, direction.

  2. La règle de la chaîne devient limpide. D(g∘f)(a)=Dg(f(a))∘Df(a)D(g \circ f)(\mathbf{a}) = Dg(f(\mathbf{a})) \circ Df(\mathbf{a}): c’est une composition de fonctions. Pas de Σ\Sigma, pas d’indices, pas de conventions d’appariement.

  3. 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 SpivakSignificationÉcriture Leibniz équivalente
Df(a)Df(\mathbf{a})Dérivée totale de ff en a\mathbf{a} (application linéaire)Jf(a)\mathbf{J}_f(\mathbf{a}) (matrice jacobienne)
Df(a)(v)Df(\mathbf{a})(\mathbf{v})Appliquer la dérivée au tangent v\mathbf{v}Jf(a) v\mathbf{J}_f(\mathbf{a})\,\mathbf{v} (JVP)
[Df(a)]∗(u)[Df(\mathbf{a})]^*(\mathbf{u})Adjoint appliqué au cotangent u\mathbf{u}Jf(a)⊤u\mathbf{J}_f(\mathbf{a})^\top\mathbf{u} (VJP)
D(g∘f)(a)D(g \circ f)(\mathbf{a})Dérivée de la compositionJg(f(a)) Jf(a)\mathbf{J}_g(f(\mathbf{a}))\,\mathbf{J}_f(\mathbf{a})

La perspective opérateur

DD est un opérateur qui transforme une fonction en une nouvelle fonction:

f  ⟼  Dff \;\longmapsto\; Df

Ensuite, Df(a)Df(\mathbf{a}) est l’application linéaire obtenue en évaluant DfDf au point a\mathbf{a}. Enfin, Df(a)(v)Df(\mathbf{a})(\mathbf{v}) est cette application linéaire appliquée au vecteur v\mathbf{v}. Trois niveaux d’« application de fonction » — et chacun correspond à un appel JAX:

NiveauNotationJAX
Opérateur de dérivationf↦Dff \mapsto Dfjax.jacobian, jax.jvp, jax.vjp
Évaluation au point a\mathbf{a}Df(a)Df(\mathbf{a})jax.jacobian(f)(a)
Application au vecteur v\mathbf{v}Df(a)(v)Df(\mathbf{a})(\mathbf{v})jax.jvp(f, (a,), (v,))[1]

Partie 2: La dérivée totale comme application linéaire

En calcul à une variable, on écrit f′(a)∈Rf'(a) \in \mathbb{R}: 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 f:Rn→Rmf: \mathbb{R}^n \to \mathbb{R}^m. La dérivée totale de ff en a\mathbf{a} est l’unique application linéaire Df(a):Rn→RmDf(\mathbf{a}): \mathbb{R}^n \to \mathbb{R}^m telle que

lim⁡h→0∥f(a+h)−f(a)−Df(a)(h)∥∥h∥=0\lim_{\mathbf{h} \to \mathbf{0}} \frac{\|f(\mathbf{a} + \mathbf{h}) - f(\mathbf{a}) - Df(\mathbf{a})(\mathbf{h})\|}{\|\mathbf{h}\|} = 0

Trois points à retenir:

  1. Df(a)Df(\mathbf{a}) est une fonction — elle prend un vecteur v∈Rn\mathbf{v} \in \mathbb{R}^n et renvoie un vecteur Df(a)(v)∈RmDf(\mathbf{a})(\mathbf{v}) \in \mathbb{R}^m.

  2. La matrice jacobienne Jf(a)∈Rm×n\mathbf{J}_f(\mathbf{a}) \in \mathbb{R}^{m \times n} représente cette application linéaire dans la base canonique: Df(a)(v)=Jf(a) vDf(\mathbf{a})(\mathbf{v}) = \mathbf{J}_f(\mathbf{a}) \, \mathbf{v}.

  3. L’écriture Df(a)(v)Df(\mathbf{a})(\mathbf{v}) est exactement ce que JAX appelle un JVP (Jacobian-vector product).

Travaillons avec un exemple concret tout au long du TP.

Exemple fil conducteur. Considérons f:R2→R3f: \mathbb{R}^2 \to \mathbb{R}^3 définie par

f(x1,x2)=(x12+x2x1x2sin⁡(x1))f(x_1, x_2) = \begin{pmatrix} x_1^2 + x_2 \\ x_1 x_2 \\ \sin(x_1) \end{pmatrix}

évaluée au point a=(1,2)\mathbf{a} = (1, 2). La jacobienne vaut

Jf(a)=(2x11x2x1cos⁡(x1)0)∣a=(1,2)=(2121cos⁡(1)0)\mathbf{J}_f(\mathbf{a}) = \begin{pmatrix} 2x_1 & 1 \\ x_2 & x_1 \\ \cos(x_1) & 0 \end{pmatrix}\bigg|_{\mathbf{a}=(1,2)} = \begin{pmatrix} 2 & 1 \\ 2 & 1 \\ \cos(1) & 0 \end{pmatrix}
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 Df(a)Df(\mathbf{a}) en tant que fonction à un vecteur tangent v=(0,5,  −1)\mathbf{v} = (0{,}5,\; -1). On peut le faire de deux façons:

  1. Produit matrice-vecteur: Jf(a) v\mathbf{J}_f(\mathbf{a}) \, \mathbf{v} (forme la matrice, puis multiplie)

  2. jax.jvp: calcule Df(a)(v)Df(\mathbf{a})(\mathbf{v}) 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 ∥h∥\|\mathbf{h}\|. En posant h=t v\mathbf{h} = t \, \mathbf{v} pour un v\mathbf{v} fixé, le rapport

r(t)=∥f(a+t v)−f(a)−Df(a)(t v)∥∥t v∥r(t) = \frac{\|f(\mathbf{a} + t\,\mathbf{v}) - f(\mathbf{a}) - Df(\mathbf{a})(t\,\mathbf{v})\|}{\|t\,\mathbf{v}\|}

doit tendre vers 0 quand t→0t \to 0. Plus précisément, on s’attend à r(t)=O(t)r(t) = O(t) 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 r(t)r(t) 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 ∂fi∂xj(a)\frac{\partial f_i}{\partial x_j}(\mathbf{a}) comme si c’était une quantité primitive. En réalité, cette dérivée partielle est une restriction de la dérivée totale Df(a)Df(\mathbf{a}). Deux opérations la définissent:

  • Injection ιj:R→Rn\iota_j: \mathbb{R} \to \mathbb{R}^n, qui envoie un scalaire tt sur le vecteur t ejt \, \mathbf{e}_j (perturbation le long de l’axe jj uniquement)

  • Projection πi:Rm→R\pi_i: \mathbb{R}^m \to \mathbb{R}, qui extrait la composante ii: πi(y)=yi\pi_i(\mathbf{y}) = y_i

La dérivée partielle est alors la composition

∂fi∂xj(a)=(πi∘Df(a)∘ιj)(1)=(Df(a)(ej))i=(Jf(a))ij\frac{\partial f_i}{\partial x_j}(\mathbf{a}) = (\pi_i \circ Df(\mathbf{a}) \circ \iota_j)(1) = \bigl(Df(\mathbf{a})(\mathbf{e}_j)\bigr)_i = \bigl(\mathbf{J}_f(\mathbf{a})\bigr)_{ij}

Autrement dit: on injecte un vecteur unitaire ej\mathbf{e}_j dans la dérivée totale, puis on projette sur la composante ii. Les entrées de la jacobienne ne sont pas la définition — elles en sont une conséquence.

En code, ιj(1)=ej\iota_j(1) = \mathbf{e}_j est simplement un vecteur canonique, et πi\pi_i 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 Df(a)(ej)Df(\mathbf{a})(\mathbf{e}_j) donne la jj-ième colonne de Jf(a)\mathbf{J}_f(\mathbf{a}): c’est le vecteur des réponses de toutes les sorties à une perturbation unitaire de l’entrée jj.

Reconstruisez la jacobienne Jf(a)\mathbf{J}_f(\mathbf{a}) colonne par colonne en utilisant n=2n = 2 appels à jax.jvp, un par vecteur canonique ej\mathbf{e}_j.

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 Jf(a)⊤u\mathbf{J}_f(\mathbf{a})^\top \mathbf{u}.

Si on choisit u=ei\mathbf{u} = \mathbf{e}_i (vecteur canonique de Rm\mathbb{R}^m), on obtient

Jf(a)⊤ei=la i-ieˋme ligne de Jf(a)\mathbf{J}_f(\mathbf{a})^\top \mathbf{e}_i = \text{la } i\text{-ième ligne de } \mathbf{J}_f(\mathbf{a})

Reconstruisez la jacobienne ligne par ligne en utilisant m=3m = 3 appels VJP.

Observation clé: l’exercice 2 utilise nn JVP (un par entrée), celui-ci utilise mm VJP (un par sortie). Lequel est plus économique dépend de nn et mm.

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

D(g∘f)(a)=Dg(f(a))∘Df(a)D(g \circ f)(\mathbf{a}) = Dg\bigl(f(\mathbf{a})\bigr) \circ Df(\mathbf{a})

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:

Jg∘f(a)=Jg(f(a))⋅Jf(a)\mathbf{J}_{g \circ f}(\mathbf{a}) = \mathbf{J}_g\bigl(f(\mathbf{a})\bigr) \cdot \mathbf{J}_f(\mathbf{a})

Définissons une seconde fonction pour composer avec ff:

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 Jh(a)\mathbf{J}_h(\mathbf{a}) de trois manières et vérifiez qu’elles coïncident:

  1. Directement: jax.jacobian(h)(a)

  2. Produit de jacobiennes: Jg(f(a))⋅Jf(a)\mathbf{J}_g(f(\mathbf{a})) \cdot \mathbf{J}_f(\mathbf{a})

  3. Composition de JVP: pour chaque ej\mathbf{e}_j, calculer Dg(f(a))(Df(a)(ej))Dg(f(\mathbf{a}))\bigl(Df(\mathbf{a})(\mathbf{e}_j)\bigr) 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 Df(a)Df(\mathbf{a}) est une application linéaire représentée par la matrice jacobienne Jf(a)∈Rm×n\mathbf{J}_f(\mathbf{a}) \in \mathbb{R}^{m \times n}. Le JVP multiplie cette matrice par un vecteur à droite: Jv\mathbf{J} \mathbf{v}.

Le VJP (vector-Jacobian product) effectue l’opération duale: on prend un vecteur u∈Rm\mathbf{u} \in \mathbb{R}^m et on le multiplie à gauche de la jacobienne:

VJP=u⊤Jf(a)\text{VJP} = \mathbf{u}^\top \mathbf{J}_f(\mathbf{a})

Le résultat est un vecteur de Rn\mathbb{R}^n, de la même taille que l’entrée de ff. Le vecteur u\mathbf{u} 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.

DirectionOpérationRôle
JVPJ v\mathbf{J} \, \mathbf{v}Propage un signal vers l’avant (entrée → sortie)
VJPu⊤J\mathbf{u}^\top \mathbf{J}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 u⊤Jf(a)\mathbf{u}^\top \mathbf{J}_f(\mathbf{a}).

Dans certains manuels, le VJP est écrit J⊤u\mathbf{J}^\top \mathbf{u} (transposer la matrice puis multiplier à droite) plutôt que u⊤J\mathbf{u}^\top \mathbf{J} (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 Df(a)Df(\mathbf{a}), noté [Df(a)]∗[Df(\mathbf{a})]^*.

# 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 ff au point a\mathbf{a} est le produit u⊤Jf(a)\mathbf{u}^\top \mathbf{J}_f(\mathbf{a}). On pourrait calculer J\mathbf{J} en entier puis effectuer le produit matriciel, mais la jacobienne est souvent trop grande pour être stockée. On cherche donc à écrire u⊤J\mathbf{u}^\top \mathbf{J} 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 f:Rn→Rmf: \mathbb{R}^n \to \mathbb{R}^m:

  1. Calculer les dérivées partielles ∂fi∂xj(a)\frac{\partial f_i}{\partial x_j}(\mathbf{a}) et les organiser dans la matrice J∈Rm×n\mathbf{J} \in \mathbb{R}^{m \times n}.

  2. Calculer u⊤J\mathbf{u}^\top \mathbf{J}. La jj-ième composante du résultat est ∑iui∂fi∂xj(a)\displaystyle\sum_i u_i \frac{\partial f_i}{\partial x_j}(\mathbf{a}).

  3. 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 f(x)=exf(\mathbf{x}) = e^{\mathbf{x}}, c’est-à-dire fi(x)=exif_i(\mathbf{x}) = e^{x_i}. Prenons a=(1,2,0)\mathbf{a} = (1, 2, 0) et u=(0,5,  −1,  2)\mathbf{u} = (0{,}5,\; -1,\; 2).

Étape 1. Chaque sortie fif_i ne dépend que de xix_i, donc les dérivées hors-diagonale sont nulles:

J=(e1000e2000e0)≈(2,720007,390001)\mathbf{J} = \begin{pmatrix} e^1 & 0 & 0 \\ 0 & e^2 & 0 \\ 0 & 0 & e^0 \end{pmatrix} \approx \begin{pmatrix} 2{,}72 & 0 & 0 \\ 0 & 7{,}39 & 0 \\ 0 & 0 & 1 \end{pmatrix}

Étape 2. On multiplie le vecteur ligne u⊤\mathbf{u}^\top par chaque colonne de J\mathbf{J}. Puisque J\mathbf{J} est diagonale, chaque colonne n’a qu’une entrée non nulle:

u⊤J=(  0,5⋅e1,    (−1)⋅e2,    2⋅e0  )≈(1,36,    −7,39,    2)\mathbf{u}^\top \mathbf{J} = \bigl(\; 0{,}5 \cdot e^1,\;\; (-1) \cdot e^2,\;\; 2 \cdot e^0 \;\bigr) \approx (1{,}36,\;\; -7{,}39,\;\; 2)

Étape 3. Chaque composante du résultat est uj⋅eaju_j \cdot e^{a_j}, ce qui donne une multiplication élément par élément:

VJP=u⊙ea\boxed{\text{VJP} = \mathbf{u} \odot e^{\mathbf{a}}}

Le coût est O(n)O(n) au lieu de O(n2)O(n^2) si on avait formé la matrice entière.


Exemple B: somme (réduction)

Soit f(x)=∑ixif(\mathbf{x}) = \sum_i x_i, qui prend un vecteur de nn nombres et renvoie un scalaire. Prenons a=(3,−1,4)\mathbf{a} = (3, -1, 4) et u=5u = 5. Ici u\mathbf{u} est un scalaire puisque la sortie l’est aussi.

Étape 1. La sortie est un scalaire, donc la jacobienne est un vecteur ligne:

J=(111)\mathbf{J} = \begin{pmatrix} 1 & 1 & 1 \end{pmatrix}

Étape 2. Un scalaire fois un vecteur ligne:

5⋅(111)=(555)5 \cdot \begin{pmatrix} 1 & 1 & 1 \end{pmatrix} = \begin{pmatrix} 5 & 5 & 5 \end{pmatrix}

Étape 3. Le scalaire uu est copié sur chaque composante:

VJP=u⋅1\boxed{\text{VJP} = u \cdot \mathbf{1}}

La somme contracte nn valeurs en une seule; son VJP fait l’inverse et diffuse le scalaire uu en nn 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: φ(x)=x⊙2\varphi(\mathbf{x}) = \mathbf{x}^{\odot 2}

La jacobienne est Dφ(a)=diag⁡(2a)D\varphi(\mathbf{a}) = \operatorname{diag}(2\mathbf{a}). La règle VJP:

[Dφ(a)]∗(u)=2a⊙u[D\varphi(\mathbf{a})]^*(\mathbf{u}) = 2\mathbf{a} \odot \mathbf{u}
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: f(z)=Wzf(\mathbf{z}) = W\mathbf{z} avec W∈Rm×nW \in \mathbb{R}^{m \times n} fixe.

La jacobienne par rapport à z\mathbf{z} est Df(a)=WDf(\mathbf{a}) = W (application linéaire constante). La règle VJP:

[Df(a)]∗(u)=W⊤u[Df(\mathbf{a})]^*(\mathbf{u}) = W^\top \mathbf{u}
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: σ:RK→RK\sigma: \mathbb{R}^K \to \mathbb{R}^K, σi(a)=eai∑jeaj\sigma_i(\mathbf{a}) = \frac{e^{a_i}}{\sum_j e^{a_j}}.

La jacobienne a pour entrées ∂σi∂aj=σi(δij−σj)\frac{\partial \sigma_i}{\partial a_j} = \sigma_i(\delta_{ij} - \sigma_j), ce qui donne la règle VJP:

[Dσ(a)]∗(u)=σ(a)⊙(u−⟨σ(a),u⟩ 1)[D\sigma(\mathbf{a})]^*(\mathbf{u}) = \sigma(\mathbf{a}) \odot \bigl(\mathbf{u} - \langle \sigma(\mathbf{a}), \mathbf{u} \rangle \, \mathbf{1}\bigr)

Ce concept est subtil. Prenez le temps de dérouler le produit Jσ⊤u\mathbf{J}_\sigma^\top \mathbf{u} à 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}:

  1. \frac{\partial (x_i^2)}{\partial x_j} = 2a_i \delta_{ij}, donc \mathbf{J} = \operatorname{diag}(2a_1, \ldots, 2a_n).

  2. (\mathbf{u}^\top \mathbf{J})_j = u_j \cdot 2a_j (matrice diagonale: chaque colonne n’a qu’une entrée).

  3. \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}:

  1. La fonction est linéaire, donc \mathbf{J} = W \in \mathbb{R}^{m \times n}.

  2. \mathbf{u}^\top W, un produit vecteur-matrice.

  3. \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}}:

  1. \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.

  2. (\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).

  3. \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 ℓ=hk∘⋯∘h1\ell = h_k \circ \cdots \circ h_1:

  1. Passe avant: identifier les valeurs intermédiaires z0=a\mathbf{z}_0 = \mathbf{a}, z1=h1(z0)\mathbf{z}_1 = h_1(\mathbf{z}_0), ..., zk=ℓ\mathbf{z}_k = \ell.

  2. Passe arrière: partir de zˉk=1\bar{\mathbf{z}}_k = 1 (cotangent initial pour une perte scalaire) et propager: zˉi−1=[Dhi(zi−1)]∗(zˉi)\bar{\mathbf{z}}_{i-1} = [Dh_i(\mathbf{z}_{i-1})]^*(\bar{\mathbf{z}}_i).

  3. Le gradient est ∇ℓ(a)=zˉ0\nabla \ell(\mathbf{a}) = \bar{\mathbf{z}}_0.

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é: ℓ(x)=∑ixi2\ell(\mathbf{x}) = \sum_i x_i^2

Décomposition: x→(⋅)2x⊙2→∑ℓ\mathbf{x} \xrightarrow{(\cdot)^2} \mathbf{x}^{\odot 2} \xrightarrow{\sum} \ell

(b) ★★ Log-sum-exp: ℓ(x)=log⁡(∑iexi)\ell(\mathbf{x}) = \log\bigl(\sum_i e^{x_i}\bigr)

Décomposition: x→exp⁡y→∑S→log⁡ℓ\mathbf{x} \xrightarrow{\exp} \mathbf{y} \xrightarrow{\sum} S \xrightarrow{\log} \ell

(c) ★★ Moindres carrés: ℓ(x)=∥Ax−b∥2\ell(\mathbf{x}) = \|A\mathbf{x} - \mathbf{b}\|^2 avec A∈Rm×nA \in \mathbb{R}^{m \times n} et b∈Rm\mathbf{b} \in \mathbb{R}^m fixes.

Décomposition: x→A⋅z→−br→(⋅)2r⊙2→∑ℓ\mathbf{x} \xrightarrow{A \cdot} \mathbf{z} \xrightarrow{- \mathbf{b}} \mathbf{r} \xrightarrow{(\cdot)^2} \mathbf{r}^{\odot 2} \xrightarrow{\sum} \ell

(d) ★★★ Entropie croisée: ℓ(z)=−log⁡(σ(z)k)\ell(\mathbf{z}) = -\log\bigl(\sigma(\mathbf{z})_k\bigr) avec kk la classe cible et σ\sigma le softmax.

Décomposition: z→σσ→πkσk→−log⁡ℓ\mathbf{z} \xrightarrow{\sigma} \boldsymbol{\sigma} \xrightarrow{\pi_k} \sigma_k \xrightarrow{-\log} \ell

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_k

Dérivation (a) — \ell(\mathbf{x}) = \sum_i x_i^2:

\mathbf{x} \xrightarrow{(\cdot)^2} \underbrace{\mathbf{x}^{\odot 2}}_{\mathbf{y}} \xrightarrow{\sum} \ell

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}):

\mathbf{x} \xrightarrow{\exp} \underbrace{e^{\mathbf{x}}}_{\mathbf{y}} \xrightarrow{\sum} \underbrace{S}_{= \sum_i e^{x_i}} \xrightarrow{\log} \ell

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:

\mathbf{x} \xrightarrow{A \cdot} \underbrace{A\mathbf{x}}_{\mathbf{z}} \xrightarrow{- \mathbf{b}} \underbrace{A\mathbf{x} - \mathbf{b}}_{\mathbf{r}} \xrightarrow{(\cdot)^2} \mathbf{r}^{\odot 2} \xrightarrow{\sum} \ell

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):

\mathbf{z} \xrightarrow{\sigma} \underbrace{\boldsymbol{\sigma}}_{\text{softmax}} \xrightarrow{\pi_k} \underbrace{\sigma_k}_{\text{scalaire}} \xrightarrow{-\log} \ell

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 D(g∘f)(a)=Dg(f(a))∘Df(a)D(g \circ f)(\mathbf{a}) = Dg(f(\mathbf{a})) \circ Df(\mathbf{a}), et en utilisant (AB)⊤=B⊤A⊤(AB)^\top = B^\top A^\top, on obtient

[D(g∘f)(a)]∗(u)=[Df(a)]∗([Dg(f(a))]∗(u))[D(g \circ f)(\mathbf{a})]^*(\mathbf{u}) = [Df(\mathbf{a})]^*\bigl([Dg(f(\mathbf{a}))]^*(\mathbf{u})\bigr)

L’ordre s’inverse: on applique d’abord l’adjoint de la fonction extérieure gg, puis celui de la fonction intérieure ff. C’est la rétropropagation: le signal u\mathbf{u} 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 ej\mathbf{e}_j donne la colonne jj de la jacobienne, un VJP avec ei\mathbf{e}_i donne la ligne ii. Pour reconstruire la jacobienne complète J∈Rm×n\mathbf{J} \in \mathbb{R}^{m \times n}, il faut donc:

  • nn passes JVP (une par colonne), ou

  • mm 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:

Fonctionnn (entrées)mm (sorties)Passes JVPPasses VJPMode le plus économique
f:R2→R3f: \mathbb{R}^2 \to \mathbb{R}^32323JVP (mode avant)
Perte scalaire: ℓ:R5→R\ell: \mathbb{R}^5 \to \mathbb{R}5151VJP (mode arrière)
Réseau de neurones: ℓ:R106→R\ell: \mathbb{R}^{10^6} \to \mathbb{R}10611061VJP (mode arrière)
Couche cachée: f:R1→R100f: \mathbb{R}^1 \to \mathbb{R}^{100}11001100JVP (mode avant)

Règle générale: le mode arrière (VJP) domine quand m≪nm \ll n. En apprentissage machine, la perte est toujours scalaire (m=1m = 1) et le nombre de paramètres est grand (n≫1n \gg 1). C’est pourquoi la rétropropagation — qui est une composition de VJP — est l’algorithme standard.

Exercice 7: Composer des VJP ★★

En reprenant h=g∘fh = g \circ f et le cotangent u=(1,−1)\mathbf{u} = (1, -1) (dimension de sortie de hh), calculez [Dh(a)]∗(u)[Dh(\mathbf{a})]^*(\mathbf{u}) 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ée

Notez 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):

f1:R100→R50,f2:R50→R20,f3:R20→R1f_1: \mathbb{R}^{100} \to \mathbb{R}^{50}, \quad f_2: \mathbb{R}^{50} \to \mathbb{R}^{20}, \quad f_3: \mathbb{R}^{20} \to \mathbb{R}^1

La composition est ℓ=f3∘f2∘f1:R100→R1\ell = f_3 \circ f_2 \circ f_1: \mathbb{R}^{100} \to \mathbb{R}^1.

Les jacobiennes sont: J1∈R50×100\mathbf{J}_1 \in \mathbb{R}^{50 \times 100}, J2∈R20×50\mathbf{J}_2 \in \mathbb{R}^{20 \times 50}, J3∈R1×20\mathbf{J}_3 \in \mathbb{R}^{1 \times 20}.

Le gradient complet est ∇ℓ=J1⊤J2⊤J3⊤∈R100\nabla \ell = \mathbf{J}_1^\top \mathbf{J}_2^\top \mathbf{J}_3^\top \in \mathbb{R}^{100} (un vecteur).

(a) Le coût d’un produit matrice-vecteur AvA \mathbf{v} avec A∈Rp×qA \in \mathbb{R}^{p \times q} est p×qp \times q multiplications. Comptez les multiplications pour les deux ordres d’évaluation:

OrdreOpérationsMultiplications
Mode avant (JVP): J3(J2(J1v))\mathbf{J}_3(\mathbf{J}_2(\mathbf{J}_1 \mathbf{v}))J1v\mathbf{J}_1 \mathbf{v}, puis J2(⋅)\mathbf{J}_2(\cdot), puis J3(⋅)\mathbf{J}_3(\cdot)? + ? + ? = ? par tangent
Mode arrière (VJP): J1⊤(J2⊤(J3⊤u))\mathbf{J}_1^\top(\mathbf{J}_2^\top(\mathbf{J}_3^\top u))J3⊤u\mathbf{J}_3^\top u, puis J2⊤(⋅)\mathbf{J}_2^\top(\cdot), puis J1⊤(⋅)\mathbf{J}_1^\top(\cdot)? + ? + ? = ? par cotangent

Pour le gradient complet (∇ℓ∈R100\nabla \ell \in \mathbb{R}^{100}): le mode avant requiert n=100n = 100 passes JVP (une par ej\mathbf{e}_j), le mode arrière requiert m=1m = 1 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: f1:R1→R20f_1: \mathbb{R}^1 \to \mathbb{R}^{20}, f2:R20→R50f_2: \mathbb{R}^{20} \to \mathbb{R}^{50}, f3:R50→R100f_3: \mathbb{R}^{50} \to \mathbb{R}^{100}. 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)
OrdreOpérationsMultiplications
Mode avant (JVP)50 \times 100 + 20 \times 50 + 1 \times 205000 + 1000 + 20 = 6020 par tangent
Mode arrière (VJP)20 \times 1 + 50 \times 20 + 100 \times 5020 + 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    # = 6020

Partie 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:

∂fi∂xj(a)≈fi(a+ε ej)−fi(a−ε ej)2ε\frac{\partial f_i}{\partial x_j}(\mathbf{a}) \approx \frac{f_i(\mathbf{a} + \varepsilon \, \mathbf{e}_j) - f_i(\mathbf{a} - \varepsilon \, \mathbf{e}_j)}{2\varepsilon}

donne une erreur O(ε2)O(\varepsilon^2) (bien meilleure que la différence avant, qui est O(ε)O(\varepsilon)). Le choix de ε\varepsilon est un compromis: trop grand → erreur de troncature; trop petit → erreurs d’arrondi en virgule flottante. En pratique, ε≈10−7\varepsilon \approx 10^{-7} 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 ff au point a\mathbf{a}.

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 J

Traçons l’erreur en fonction de ε\varepsilon 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()
<Figure size 800x500 with 1 Axes>

Récapitulatif

Ce TP a relié la notation de Spivak au code JAX. Le tableau suivant résume les correspondances.

Concept (Spivak)NotationCode JAXInterprétation
Dérivée totaleDf(a)Df(\mathbf{a})jax.jacobian(f)(a)Application linéaire Rn→Rm\mathbb{R}^n \to \mathbb{R}^m
Application à un tangentDf(a)(v)Df(\mathbf{a})(\mathbf{v})jax.jvp(f, (a,), (v,))[1]JVP, mode avant
VJPu⊤Jf(a)\mathbf{u}^\top \mathbf{J}_f(\mathbf{a})jax.vjp(f, a)[1](u)[0]VJP, mode arrière
Dérivée partielleπi∘Df(a)∘ιj\pi_i \circ Df(\mathbf{a}) \circ \iota_jjax.jacobian(f)(a)[i, j]Restriction de Df(a)Df(\mathbf{a})
Règle de la chaîneD(g∘f)(a)=Dg(f(a))∘Df(a)D(g \circ f)(\mathbf{a}) = Dg(f(\mathbf{a})) \circ Df(\mathbf{a})Composer jax.jvp ou jax.vjpComposition d’opérateurs linéaires

À retenir:

  1. La dérivée totale est une application linéaire, pas une matrice. La matrice jacobienne la représente.

  2. Les dérivées partielles sont des restrictions: ∂fi∂xj=πi∘Df(a)∘ιj\frac{\partial f_i}{\partial x_j} = \pi_i \circ Df(\mathbf{a}) \circ \iota_j.

  3. La règle de la chaîne est une composition d’applications linéaires. Pour le VJP, l’ordre de composition s’inverse.

  4. Le VJP est le produit vecteur-jacobienne: u⊤Jf(a)\mathbf{u}^\top \mathbf{J}_f(\mathbf{a}).

  5. Pour f:Rn→R1f: \mathbb{R}^n \to \mathbb{R}^1 (perte scalaire), un seul VJP donne le gradient complet, contre nn 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 ».