Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions .translate/state/numpy_vs_numba_vs_jax.md.yml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
source-sha: d08a73d48a409509d7d6f6585b99c2c8909c9a28
synced-at: "2026-07-14"
model: claude-opus-4-8
mode: NEW
source-sha: c2589a2b61df5c841753b7a9018fda63ad0711fc
synced-at: "2026-08-20"
model: claude-sonnet-5
mode: UPDATE
section-count: 3
tool-version: 0.15.0
tool-version: 0.26.0
17 changes: 2 additions & 15 deletions lectures/numpy_vs_numba_vs_jax.md
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,7 @@ En plus de ce qui est inclus dans Anaconda, ce cours nécessitera les bibliothè
---
tags: [hide-output]
---
!pip install quantecon jax
!pip install quantecon "jax==0.11.0"
```

```{include} _admonition/gpu.md
Expand Down Expand Up @@ -143,7 +143,6 @@ for x in grid:
m = max(m, z)
```


### Vectorisation avec NumPy

Passons à NumPy et utilisons une grille plus grande
Expand Down Expand Up @@ -208,7 +207,6 @@ De plus, l'exécution eager de NumPy crée de nombreux tableaux intermédiaires

Ce type d'utilisation de la mémoire peut poser un gros problème dans les calculs de recherche réels.


### Une comparaison avec Numba

Voyons si nous pouvons obtenir de meilleures performances en utilisant Numba avec une simple boucle.
Expand Down Expand Up @@ -250,7 +248,6 @@ Sur la plupart des machines, la version Numba sera un peu plus rapide que NumPy.

La raison en est un code machine efficace ainsi que moins d'opérations de lecture-écriture en mémoire.


### Numba parallélisé

Essayons maintenant la parallélisation avec Numba en utilisant `prange` :
Expand Down Expand Up @@ -296,7 +293,6 @@ print(f"Numba result: {z_max_parallel:.6f}")

Pour les machines puissantes et les grilles de plus grande taille, la parallélisation peut générer des gains de vitesse utiles, même sur le CPU.


### Code vectorisé avec JAX

Essayons de reproduire l'approche vectorisée de NumPy avec JAX.
Expand Down Expand Up @@ -343,7 +339,6 @@ Une fois compilé, JAX est nettement plus rapide que NumPy, en particulier sur u

Le surcoût de compilation est un coût ponctuel qui est rentabilisé lorsque la fonction est appelée à plusieurs reprises.


### JAX plus vmap

Comme nous avons utilisé `jax.jit` ci-dessus, nous avons évité de créer de nombreux tableaux intermédiaires.
Expand Down Expand Up @@ -400,7 +395,6 @@ with qe.Timer():
z_max.block_until_ready()
```


### Résumé

À notre avis, JAX est le gagnant pour les opérations vectorisées.
Expand All @@ -413,7 +407,6 @@ Il domine également Numba lorsqu'il est exécuté sur le GPU.
Numba peut prendre en charge la programmation GPU via `numba.cuda`, mais nous devons alors paralléliser à la main. Pour la plupart des cas rencontrés en économie, en économétrie et en finance, il est bien préférable de laisser le compilateur JAX gérer une parallélisation efficace plutôt que d'essayer de coder ces routines nous-mêmes.
```


## Opérations séquentielles

Certaines opérations sont intrinsèquement séquentielles -- et donc difficiles voire impossibles à vectoriser.
Expand All @@ -422,7 +415,6 @@ Dans ce cas, NumPy est une mauvaise option et il ne nous reste que le choix entr

Pour comparer ces choix, nous reviendrons sur le problème de l'itération sur l'application quadratique que nous avons vu dans notre {doc}`cours sur Numba <numba>`.


### Version Numba

Voici la version Numba.
Expand Down Expand Up @@ -457,7 +449,6 @@ with qe.Timer():

Numba gère cette opération séquentielle de manière très efficace.


### Version JAX

Nous ne pouvons pas remplacer directement `numba.jit` par `jax.jit` car les tableaux JAX sont immuables.
Expand Down Expand Up @@ -513,7 +504,6 @@ with qe.Timer():

JAX est également assez efficace pour cette opération séquentielle !


#### Deuxième tentative

Il existe une autre manière d'implémenter la boucle qui utilise `lax.scan`.
Expand Down Expand Up @@ -556,7 +546,6 @@ with qe.Timer():

Étonnamment, JAX offre également de solides performances après compilation.


### Résumé

Bien que Numba et JAX offrent tous deux de solides performances pour les opérations séquentielles, il existe des différences en termes de lisibilité du code et de facilité d'utilisation.
Expand All @@ -569,8 +558,6 @@ Les versions JAX, en revanche, nécessitent soit `lax.fori_loop`, soit `lax.scan

Bien que la syntaxe `at[t].set` de JAX permette des mises à jour élément par élément, le code global reste plus difficile à lire que l'équivalent Numba.



## Recommandations générales

Prenons maintenant du recul et résumons les compromis.
Expand All @@ -591,4 +578,4 @@ JAX peut gérer les problèmes séquentiels via `lax.fori_loop` ou `lax.scan`, m

D'un autre côté, les versions JAX prennent en charge la différenciation automatique.

Cela pourrait présenter un intérêt si, par exemple, nous souhaitons calculer les sensibilités d'une trajectoire aux paramètres du modèle
Cela pourrait présenter un intérêt si, par exemple, nous souhaitons calculer les sensibilités d'une trajectoire aux paramètres du modèle
Loading