En esta guía, se explica cómo usar el backend de JAX en Meridian.
Introducción al backend de JAX
Meridian usa JAX como su backend numérico predeterminado para las operaciones numéricas fundamentales y el muestreo probabilístico de Markov Chain Monte Carlo (MCMC) (a partir de Meridian 2.0). JAX fomenta un estilo de programación funcional y utiliza la compilación de XLA (Accelerated Linear Algebra) para ofrecer optimizaciones avanzadas del rendimiento y eficiencia de la memoria.
El backend heredado de TensorFlow dejó de estar disponible y se quitará en una versión futura.
Instructivo: Para ver JAX en acción, consulta el notebook Primeros pasos con JAX.
Configuración del backend
De forma predeterminada, Meridian se ejecuta en JAX. No es necesario configurar ninguna variable de entorno para usar JAX.
Backend de TensorFlow heredado (obsoleto)
Si necesitas ejecutar temporalmente con el backend de TensorFlow heredado, configura la variable de entorno MERIDIAN_BACKEND como 'tensorflow' antes de importar 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
Configuración de la precisión
De forma predeterminada, Meridian se ejecuta con una precisión de 64 bits (float64) en JAX.
En cambio, los usuarios pueden optar por ejecutar con una precisión de 32 bits (float32), por ejemplo:
- Tiempo de ejecución de entrenamiento más rápido: Las operaciones de punto flotante de 32 bits se pueden ejecutar más rápido en aceleradores de hardware (como GPU o TPU).
- Menor uso de memoria: La precisión de 32 bits reduce el consumo de memoria durante el muestreo de MCMC.
Para usar la precisión de 32 bits, establece la variable de entorno MERIDIAN_ENABLE_JAX_X64 en 'False' (o '0') antes de importar 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 de entorno MERIDIAN_ENABLE_JAX_X64 no está configurada o se establece en 'True' o '1', Meridian usa la precisión de 64 bits de forma predeterminada.
Coherencia de tipos
Dado que Meridian opera con una precisión de 64 bits de forma predeterminada, asegúrate de que todos los valores proporcionados por el usuario, los arrays personalizados y los parámetros de distribución mantengan la coherencia de los tipos:
- Literales y arrays de punto flotante: Los literales de números de punto flotante estándar de Python (como
0.2,0.9) se establecen de forma predeterminada en números de punto flotante de 64 bits. Cuando crees arrays de NumPy para distribuciones a priori o entradas, usanp.float64odtype=np.float64para que coincidan con la precisión predeterminada. - Crea distribuciones a priori personalizadas con la precisión correspondiente: Cuando definas distribuciones a priori personalizadas en
PriorDistribution, asegúrate de que todos los parámetros de distribución (comoloc,scale,concentration0yconcentration1) coincidan con la precisión activa (número de punto flotante de 64 bits de forma predeterminada).
Diferencias en la API cuando se usa JAX en comparación con TensorFlow
Cuando se usa el backend de JAX, hay diferencias clave en la API que se deben tener en cuenta:
Distribuciones a priori
Los modelos de Meridian usan TensorFlow Probability en JAX (tensorflow_probability.substrates.jax). Cuando configures distribuciones a priori personalizadas en JAX, importa tensorflow_probability.substrates.jax as tfp_jax y construye distribuciones con tfp_jax.distributions.
Asegúrate de que todos los parámetros de distribución personalizada usen una precisión de 64 bits (como np.float64 o arrays de números de punto flotante de 64 bits) para mantener la coherencia de los tipos con la configuración de precisión predeterminada 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 (obsoleto)
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
)
)
Requisito de una semilla explícita
Cuando se usa el backend de JAX, se requiere una semilla explícita para las funciones estocásticas (por ejemplo, en sample_posterior()). Si bien TensorFlow usa un generador global de números aleatorios que elige automáticamente una semilla aleatoria, JAX hace que esta semilla sea explícita. No encontramos diferencias estadísticamente significativas en las estimaciones del ROI ni en los cambios de presupuesto entre las diferentes semillas.
# 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,
)
Para obtener más información sobre las semillas y los números aleatorios de JAX, consulta la documentación sobre los números pseudoaleatorios de JAX.
Diferencias numéricas y reproducibilidad
Dado que TensorFlow y JAX compilan sus grafos de procesamiento de manera diferente, es posible que observes pequeñas diferencias numéricas en tus estimaciones a posteriori cuando cambies a JAX con los mismos datos y las mismas semillas aleatorias.
Si bien las distribuciones a posteriori podrían no ser idénticas en todos los backends, las diferencias suelen ser pequeñas y no tienen importancia estadística para las métricas comerciales, como el ROI y la asignación del presupuesto. Esto garantiza que el cambio al backend de JAX mantenga la integridad de las estadísticas de tu modelo.
Consideraciones de rendimiento
Las pruebas internas revelaron que JAX potenció las ejecuciones iniciales del modelo, lo que redujo el tiempo de ejecución promedio en un 40% y el uso de memoria en un 70%, en comparación con TensorFlow cuando se usan GPU. JAX también optimizó las iteraciones del modelo, lo que permitió tiempos de ejecución 2 veces más rápidos, un uso de memoria 4 veces menor y flujos de trabajo ininterrumpidos, ya que se eliminó la necesidad de reiniciar el kernel.
Gracias a la mayor eficiencia de la memoria, tienes más margen para ajustar los parámetros que requieren una gran cantidad de procesamiento. Por ejemplo, en Meridian.sample_posterior(), puedes aumentar el argumento unrolled_leapfrog_steps (p. ej., de 1 a 5). Esto puede acelerar la convergencia, ya que aumenta la longitud de la trayectoria del No-U-Turn-Sampler (NUTS) sin exceder los límites de memoria del hardware. También puedes aumentar el parámetro n_adapt para ayudar aún más a la convergencia durante la fase de adaptación.