Ce guide explique comment utiliser le backend JAX dans Meridian.
Présentation du backend JAX
Meridian utilise JAX comme backend numérique par défaut pour les opérations numériques de base et l'échantillonnage probabiliste MCMC (Monte-Carlo par chaînes de Markov), à partir de Meridian 2.0. JAX encourage un style de programmation fonctionnel et utilise la compilation XLA (Accelerated Linear Algebra) pour des optimisations de performances avancées et une efficacité de mémoire.
L'ancien backend TensorFlow est obsolète et sera supprimé dans une prochaine version.
Tutoriel : pour voir JAX en action, consultez le notebook d'introduction à JAX.
Configuration du backend
Par défaut, Meridian s'exécute sur JAX. Vous n'avez pas besoin de configurer de variables d'environnement pour utiliser JAX.
Ancien backend TensorFlow (obsolète)
Si vous devez exécuter temporairement l'ancien backend TensorFlow, définissez la variable d'environnement MERIDIAN_BACKEND sur 'tensorflow' avant d'importer Meridian :
import os
# Select legacy TensorFlow backend (deprecated)
os.environ['MERIDIAN_BACKEND'] = 'tensorflow'
# Now it is safe to import Meridian modules
from meridian.model import model
from meridian.data import load
Configuration de précision
Par défaut, Meridian s'exécute avec une précision 64 bits (float64) sur JAX.
Les utilisateurs peuvent choisir la précision 32 bits (float32) à la place, par exemple dans les cas suivants :
- Exécution plus rapide de l'entraînement : les opérations 32 bits à virgule flottante peuvent s'exécuter plus vite sur les accélérateurs matériels (comme les GPU ou TPU).
- Utilisation de mémoire réduite : la précision 32 bits réduit la consommation de mémoire lors de l'échantillonnage MCMC.
Pour utiliser la précision 32 bits, définissez la variable d'environnement MERIDIAN_ENABLE_JAX_X64 sur 'False' (ou '0') avant d'importer Meridian :
import os
# Disable 64-bit precision (enable 32-bit precision)
os.environ['MERIDIAN_ENABLE_JAX_X64'] = 'False'
# Now it is safe to import Meridian modules
from meridian.model import model
Si la variable d'environnement MERIDIAN_ENABLE_JAX_X64 n'est pas définie ou est définie sur 'True' ou '1', Meridian utilise par défaut la précision 64 bits.
Cohérence des types
Étant donné que Meridian fonctionne par défaut avec une précision 64 bits, assurez-vous que toutes les valeurs fournies par l'utilisateur, les tableaux personnalisés et les paramètres de distribution maintiennent la cohérence des types :
- Littéraux et tableaux à virgule flottante : les littéraux à virgule flottante Python standards (tels que
0.2,0.9) sont par défaut des floats 64 bits. Lorsque vous créez des tableaux NumPy pour des a priori ou des entrées, utiliseznp.float64oudtype=np.float64pour respecter la précision par défaut. - Créer des distributions a priori personnalisées avec une précision correspondante : lorsque vous définissez des distributions a priori personnalisées dans
PriorDistribution, assurez-vous que tous les paramètres de distribution (tels queloc,scale,concentration0etconcentration1) correspondent à la précision active (float 64 bits par défaut).
Différences d'API entre JAX et TensorFlow
Lorsque vous utilisez le backend JAX, vous devez tenir compte de différences clés au niveau de l'API :
Distributions a priori
Les modèles Meridian utilisent TensorFlow Probability sur JAX (tensorflow_probability.substrates.jax). Lorsque vous configurez des distributions a priori personnalisées sous JAX, importez tensorflow_probability.substrates.jax as tfp_jax et créez des distributions à l'aide de tfp_jax.distributions.
Assurez-vous que tous les paramètres de distribution personnalisés utilisent une précision 64 bits (comme np.float64 ou des tableaux à virgule flottante 64 bits) pour maintenir la cohérence des types avec les paramètres de précision par défaut de Meridian.
JAX
import numpy as np
import tensorflow_probability.substrates.jax as tfp_jax
from meridian.model import constants
from meridian.model import prior_distribution
# Parameters use 64-bit precision
roi_mu = np.float64(0.2)
roi_sigma = np.float64(0.9)
prior = prior_distribution.PriorDistribution(
roi_m=tfp_jax.distributions.LogNormal(
roi_mu, roi_sigma, name=constants.ROI_M
)
)
TensorFlow (obsolète)
import tensorflow_probability as tfp
from meridian.model import constants
from meridian.model import prior_distribution
roi_mu = 0.2
roi_sigma = 0.9
prior = prior_distribution.PriorDistribution(
roi_m=tfp.distributions.LogNormal(
roi_mu, roi_sigma, name=constants.ROI_M
)
)
Exigence de graine explicite
Lorsque vous utilisez le backend JAX, une graine explicite est requise pour les fonctions stochastiques (par exemple, dans sample_posterior()). Alors que TensorFlow utilise un générateur de nombres aléatoires global qui choisit automatiquement une graine aléatoire, JAX rend cette graine explicite. Nous n'avons constaté aucune différence statistiquement pertinente dans les estimations du ROI ni dans les réaffectations de budget entre les différentes graines.
# Explicitly set a seed for MCMC sampling when using the JAX backend
mmm.sample_posterior(
n_chains=2,
n_adapt=1000,
n_burnin=500,
n_keep=1000,
seed=0,
)
Pour en savoir plus sur les nombres aléatoires et les graines JAX, consultez la documentation sur les nombres pseudo-aléatoires JAX.
Différences numériques et reproductibilité
Étant donné que TensorFlow et JAX compilent leurs graphes de calcul différemment, vous pouvez observer de légères différences numériques dans vos estimations a posteriori lorsque vous passez à JAX en utilisant les mêmes données et les mêmes graines aléatoires.
Bien que les distributions a posteriori puissent ne pas être identiques entre les backends, les différences sont généralement faibles et statistiquement non pertinentes pour les métriques commerciales telles que le ROI et la répartition du budget. Cela permet de s'assurer que le passage au backend JAX préserve l'intégrité des insights de votre modèle.
Considérations sur les performances
Les tests internes ont montré que JAX optimisait les exécutions initiales de modèles, réduisant le temps d'exécution moyen d'environ 40 % et l'utilisation de la mémoire d'environ 70 % par rapport à TensorFlow lors de l'utilisation de GPU. JAX a également simplifié les itérations de modèles, ce qui a permis de réduire de moitié les temps d'exécution, de diviser l'utilisation de la mémoire par quatre et de ne pas interrompre les workflows en éliminant le besoin de redémarrer le kernel.
Grâce à l'amélioration de l'efficacité de la mémoire, vous disposez d'une plus grande marge de manœuvre pour ajuster les paramètres gourmands en ressources de calcul. Ainsi, dans Meridian.sample_posterior(), vous pouvez augmenter l'argument unrolled_leapfrog_steps (par exemple, de 1 à 5). Cela peut accélérer la convergence en augmentant la longueur de la trajectoire de l'échantillonneur NUTS (No-U-Turn-Sampler) sans dépasser les limites de mémoire matérielle. Vous pouvez également augmenter le paramètre n_adapt pour faciliter davantage la convergence pendant la phase d'adaptation.