Usar o back-end JAX

Este guia explica como usar o back-end do JAX no Meridian.

Introdução ao back-end do JAX

O Meridian usa o JAX como back-end numérico padrão para as principais operações numéricas e a amostragem de Monte Carlo via cadeias de Markov (MCMC, na sigla em inglês) probabilística (a partir do Meridian 2.0). O JAX incentiva um estilo de programação funcional e usa a compilação de álgebra linear acelerada (XLA, na sigla em inglês) para oferecer otimizações avançadas de desempenho e eficiência de memória.

O back-end legado do TensorFlow foi descontinuado e será removido em uma versão futura.

Tutorial: para ver o JAX em ação, consulte o notebook Introdução ao JAX.

Configuração de back-end

Por padrão, o Meridian é executado no JAX. Não é necessário configurar nenhuma variável de ambiente para usar o JAX.

Back-end legado do TensorFlow (descontinuado)

Se você precisar executar temporariamente com o back-end legado do TensorFlow, defina a variável de ambiente MERIDIAN_BACKEND como 'tensorflow' antes de importar o 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

Configuração de precisão

Por padrão, o Meridian é executado com precisão de 64 bits (float64) no JAX.

Os usuários podem optar por executar com precisão de 32 bits (float32), por exemplo:

  • Tempo de execução de treinamento mais rápido: as operações de uso de pontos flutuantes de 32 bits podem ser executadas mais rapidamente em aceleradores de hardware (como GPUs ou TPUs).
  • Menor uso da memória: a precisão de 32 bits reduz o consumo de memória durante a amostragem de MCMC.

Para usar a precisão de 32 bits, defina a variável de ambiente MERIDIAN_ENABLE_JAX_X64 como 'False' (ou '0') antes de importar o 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

Se a variável de ambiente MERIDIAN_ENABLE_JAX_X64 não estiver definida ou estiver definida como 'True' ou '1', o Meridian vai usar a precisão de 64 bits por padrão.

Consistência de tipo

Como o Meridian opera em precisão de 64 bits por padrão, verifique se todos os valores fornecidos pelo usuário, matrizes personalizadas e parâmetros de distribuição mantêm a consistência de tipo:

  • Literais e matrizes flutuantes: os literais flutuantes padrão do Python (como 0.2, 0.9) usam como padrão pontos flutuantes de 64 bits. Ao criar matrizes NumPy para distribuições a priori ou entradas, use np.float64 ou dtype=np.float64 para corresponder à precisão padrão.
  • Crie distribuições a priori personalizadas com precisão correspondente: ao definir distribuições a priori personalizadas em PriorDistribution, verifique se todos os parâmetros de distribuição (como loc, scale, concentration0 e concentration1) correspondem à precisão ativa (ponto flutuante de 64 bits por padrão).

Diferenças de API ao usar JAX em vez de TensorFlow

Ao usar o back-end JAX, há diferenças importantes na API que você precisa considerar:

Distribuições a priori

Os modelos do Meridian usam o TensorFlow Probability no JAX (tensorflow_probability.substrates.jax). Ao configurar distribuições a priori personalizadas no JAX, importe tensorflow_probability.substrates.jax as tfp_jax e crie distribuições usando tfp_jax.distributions.

Verifique se todos os parâmetros de distribuição personalizada usam precisão de 64 bits (como np.float64 ou matrizes flutuantes de 64 bits) para manter a consistência de tipo com as configurações de precisão padrão do 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 (descontinuado)

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 semente explícita

Ao usar o back-end JAX, é necessário ter uma semente explícita para funções estocásticas (por exemplo, em sample_posterior()). Enquanto o TensorFlow usa um gerador global de números aleatórios que escolhe automaticamente uma semente, o JAX torna essa semente explícita. Não encontramos diferenças estatisticamente significativas nas estimativas de ROI ou nas mudanças de orçamento em diferentes sementes.

# 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 mais detalhes sobre números aleatórios e sementes do JAX, consulte a documentação de números pseudoaleatórios do JAX.

Diferenças numéricas e reprodutibilidade

Como o TensorFlow e o JAX compilam os grafos computacionais de maneira diferente, você pode observar pequenas diferenças numéricas nas estimativas a posteriori ao mudar para o JAX usando os mesmos dados e sementes aleatórias.

Embora as distribuições a posteriori não sejam idênticas em todos os back-ends, as diferenças geralmente são pequenas e não estatisticamente significativas para métricas de negócios como ROI e alocação de orçamento. Isso garante que a troca para o back-end do JAX mantenha a integridade dos insights do modelo.

Considerações sobre desempenho

Testes internos descobriram que o JAX impulsionou as execuções iniciais do modelo, reduzindo o tempo de execução médio em cerca de 40% e o uso da memória em cerca de 70%, em comparação com o TensorFlow ao usar GPUs. O JAX também simplificou as iterações do modelo, permitindo tempos de execução 2 vezes mais rápidos, uso de memória 4 vezes menor e fluxos de trabalho ininterruptos ao eliminar a necessidade de reinicializações do kernel.

Devido ao aumento da eficiência da memória, você tem mais espaço para ajustar parâmetros computacionalmente intensivos. Por exemplo, em Meridian.sample_posterior(), você pode aumentar o argumento unrolled_leapfrog_steps (por exemplo, de 1 para 5). Isso pode acelerar a convergência aumentando o comprimento da trajetória do No U Turn Sampler (NUTS) sem exceder os limites de memória do hardware. Também é possível aumentar o parâmetro n_adapt para ajudar ainda mais na convergência durante a fase de adaptação.