使用 JAX 后端

本指南介绍了如何在 Meridian 中使用 JAX 后端。

JAX 后端简介

自 Meridian 2.0 版本起,Meridian 已将 JAX 作为默认数值后端,用于处理核心数值运算以及概率性马尔可夫链蒙特卡洛 (MCMC) 抽样。JAX 提倡函数式编程风格,并利用 XLA(加速线性代数)编译技术来实现高级性能优化和内存效率。

旧版 TensorFlow 后端已弃用,并将在未来版本中移除。

教程:如需查看 JAX 的实际应用,请参阅开始使用 JAX 笔记本。

后端配置

默认情况下,Meridian 基于 JAX 运行。使用 JAX 无需配置任何环境变量。

旧版 TensorFlow 后端(已弃用)

如果您需要暂时使用旧版 TensorFlow 后端运行,请在导入 Meridian 之前将 MERIDIAN_BACKEND 环境变量设置为 'tensorflow'

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

精度配置

默认情况下,Meridian 基于 JAX 以 64 位精度 (float64) 运行。

用户可以选择改为以 32 位精度 (float32) 运行,例如:

  • 更快的训练运行时:32 位浮点运算在硬件加速器(例如 GPU 或 TPU)上运行速度更快。
  • 更低的内存占用:32 位精度可显著降低 MCMC 采样过程中的内存消耗。

如需使用 32 位精度,请在导入 Meridian 之前将 MERIDIAN_ENABLE_JAX_X64 环境变量设置为 'False'(或 '0'):

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

如果 MERIDIAN_ENABLE_JAX_X64 环境变量未设置或设置为 'True''1',Meridian 将默认采用 64 位精度。

类型一致性

由于 Meridian 默认以 64 位精度运行,请务必确保所有用户提供的值、自定义数组以及分布参数均保持类型一致:

  • 浮点字面量和数组:标准 Python 浮点字面量(例如 0.20.9)默认采用 64 位浮点数。为先验或输入创建 NumPy 数组时,请使用 np.float64dtype=np.float64 来匹配默认精度。
  • 构建具有匹配精度的自定义先验分布:PriorDistribution 中定义自定义先验分布时,请确保所有分布参数(例如 locscaleconcentration0concentration1)都与当前生效的精度(默认为 64 位浮点数)匹配。

JAX 与 TensorFlow 在 API 方面的差异

使用 JAX 后端时,请注意以下关键 API 差异:

先验分布

Meridian 模型使用基于 JAX 的 TensorFlow Probability (tensorflow_probability.substrates.jax)。在 JAX 下配置自定义先验分布时,请导入 tensorflow_probability.substrates.jax as tfp_jax 并使用 tfp_jax.distributions 来构建分布。

确保所有自定义分布参数都使用 64 位精度(例如 np.float64 或 64 位浮点数组),以便与 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(已弃用)

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

要求显式指定种子

使用 JAX 后端时,必须为随机函数(例如 sample_posterior() 中的函数)显式提供种子。TensorFlow 使用全局随机数生成器来自动选择随机种子,JAX 则要求显式指定种子。我们发现,不同种子生成的投资回报率估算值或预算调整幅度并无达到统计显著性的差异。

# 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,
)

如需详细了解 JAX 随机数和种子,请参阅 JAX 伪随机数文档

数值差异和可复现性

由于 TensorFlow 和 JAX 采用不同的计算图编译机制,因此在切换至 JAX 后端时,即使数据和随机种子完全一致,后验估计值仍可能存在细微的数值差异。

虽然不同后端生成的后验分布可能不尽相同,但这些差异通常微乎其微,对于投资回报率和预算分配等业务指标而言,并不具备统计显著性。这可以确保切换至 JAX 后端时,模型生成的数据洞见依然完整可靠。

性能考虑因素

内部测试发现,在使用 GPU 的环境下,相较于 TensorFlow,JAX 可大幅提升模型初始运行效率,平均运行时间缩短约 40%,内存用量降低约 70%。JAX 还优化了模型迭代流程,运行速度提升至原来的 2 倍,内存用量减少四分之三;此外,由于无需重启内核,可确保工作流不间断运行。

由于内存效率显著提升,您将拥有更大的余量来调优那些计算密集型参数。例如,在 Meridian.sample_posterior() 中,您可以调高 unrolled_leapfrog_steps 参数(例如,从 1 增至 5)。这有助于在不超出硬件内存限制的前提下,通过增加 No-U-Turn-Sampler (NUTS) 的轨迹长度来加速收敛。您还可以调高 n_adapt 参数,在自适应阶段进一步辅助模型收敛。