常见问题解答 (FAQ)#

我们在此收集常见问题的解答。欢迎贡献!

jit 改变了函数的行为#

如果你的 Python 函数在使用 jax.jit() 后行为发生了变化,可能是因为你的函数使用了全局状态或具有副作用。在下面的代码中,impure_func 使用了全局变量 y,并且由于 print 产生了副作用。

y = 0

# @jit   # Different behavior with jit
def impure_func(x):
  print("Inside:", y)
  return x + y

for y in range(3):
  print("Result:", impure_func(y))

没有 jit 时,输出为

Inside: 0
Result: 0
Inside: 1
Result: 2
Inside: 2
Result: 4

使用 jit 时,输出为

Inside: 0
Result: 0
Result: 1
Result: 2

对于 jax.jit(),函数会使用 Python 解释器执行一次,此时会触发 Inside 的打印,并观测到 y 的初始值。随后,函数被编译并缓存,之后在用不同的 x 值执行时,依然会使用最初缓存的 y 值。

延伸阅读

jit 改变了输出的精确数值#

有时用户会惊讶地发现,使用 jit() 封装函数会改变函数的输出。例如

>>> from jax import jit
>>> import jax.numpy as jnp
>>> def f(x):
...   return jnp.log(jnp.sqrt(x))
>>> x = jnp.pi
>>> print(f(x))
0.572365
>>> print(jit(f)(x))
0.5723649

这种微小的输出差异源于 XLA 编译器的优化:在编译过程中,XLA 有时会重新排列或删除某些操作,以提高整体计算效率。

在本例中,XLA 利用对数的属性将 log(sqrt(x)) 替换为 0.5 * log(x),这是一个数学上等价但计算效率更高的表达式。输出差异源于浮点运算只是真实数学的近似,因此计算同一表达式的不同方式可能会产生细微不同的结果。

其他时候,XLA 的优化可能会导致更显著的差异。考虑以下示例

>>> def f(x):
...   return jnp.log(jnp.exp(x))
>>> x = 100.0
>>> print(f(x))
inf
>>> print(jit(f)(x))
100.0

在非 JIT 编译的逐操作模式下,结果为 inf,因为 jnp.exp(x) 溢出并返回 inf。而在 JIT 下,XLA 识别出 logexp 的逆运算,从而在编译后的函数中移除了这些操作,直接返回输入值。在这种情况下,JIT 编译产生了更精确的浮点近似结果。

遗憾的是,XLA 的代数简化完整列表并未得到充分记录,但如果你熟悉 C++ 并好奇 XLA 编译器进行了哪些类型的优化,可以在源代码中查看:algebraic_simplifier.cc

jit 装饰的函数编译非常慢#

如果你的 jit 装饰函数在第一次调用时需要运行几十秒(甚至更久!),但在后续调用时执行迅速,说明 JAX 在追踪或编译你的代码时耗时较长。

这通常表明调用函数在 JAX 内部表示中生成了大量代码,典型原因是对 Python 控制流(如 for 循环)的过度使用。对于少量的循环迭代,Python 处理没问题,但如果需要 大量 循环迭代,你应该重写代码以利用 JAX 的 结构化控制流原语(如 lax.scan()),或者避免对整个循环使用 jit 装饰(你仍然可以在循环 内部 使用 jit 装饰的函数)。

如果你不确定这是否是问题所在,可以尝试对函数运行 jax.make_jaxpr()。如果输出有成百上千行,编译缓慢是预料之中的。

有时很难重写代码以避开 Python 循环,因为代码中处理了许多不同形状的数组。在这种情况下,推荐的解决方案是利用 jax.numpy.where() 等函数,在具有固定形状的填充数组上进行计算。

如果你的函数因其他原因编译缓慢,请在 GitHub 上提出问题。

如何对方法使用 jit#

已移至 🔪 在类方法中使用 jax.jit

JAX 比 NumPy 快吗?#

用户经常试图通过基准测试来回答 JAX 是否比 NumPy 快的问题;由于这两个包的差异,没有简单的答案。

概括来说

  • NumPy 操作是即时、同步执行的,且仅在 CPU 上运行。

  • JAX 操作可能是即时执行,也可能在编译后执行(如果在 jit() 内部);它们是异步调度的(参见 异步调度),并且可以在 CPU、GPU 或 TPU 上执行,这些硬件的性能特征差异巨大且在不断演变。

这些架构差异使得在 NumPy 和 JAX 之间进行有意义的直接基准测试比较变得困难。

此外,这些差异导致了两个包不同的工程重点:例如,NumPy 在降低单个数组操作的逐次调用开销方面投入了大量精力,因为在 NumPy 的计算模型中,这种开销是无法避免的。相反,JAX 有多种方式可以避免调度开销(例如 JIT 编译、异步调度、批处理变换等),因此降低逐次调用开销并非首要优先级。

综上所述:如果你在 CPU 上对单个数组操作进行微基准测试,通常 NumPy 会因其更低的逐操作调度开销而胜过 JAX。如果你在 GPU 或 TPU 上运行代码,或者在 CPU 上对更复杂的 JIT 编译操作序列进行基准测试,通常 JAX 会胜过 NumPy。

在使用 where 时,梯度包含 NaN#

如果你定义函数时使用 where 来避免未定义值,如果不小心,在反向微分时可能会得到 NaN

def my_log(x):
  return jnp.where(x > 0., jnp.log(x), 0.)

my_log(0.) ==> 0.  # Ok
jax.grad(my_log)(0.)  ==> NaN

简短的解释是,在 grad 计算期间,对应未定义的 jnp.log(x) 的伴随项 (adjoint) 是 NaN,它会被累加到 jnp.where 的伴随项中。编写此类函数的正确方法是确保部分定义的函数 内部 包含一个 jnp.where,以确保伴随项始终有限。

def safe_for_grad_log(x):
  return jnp.log(jnp.where(x > 0., x, 1.))

safe_for_grad_log(0.) ==> 0.  # Ok
jax.grad(safe_for_grad_log)(0.)  ==> 0.  # Ok

除了原始的 jnp.where,可能还需要内部的 jnp.where,例如:

def my_log_or_y(x, y):
  """Return log(x) if x > 0 or y"""
  return jnp.where(x > 0., jnp.log(jnp.where(x > 0., x, 1.)), y)

延伸阅读

为什么基于排序顺序的函数的梯度为零?#

如果你定义的函数处理输入的操作依赖于输入的相对顺序(例如 maxgreaterargsort 等),你可能会惊讶地发现梯度到处都是零。这是一个例子,我们将 f(x) 定义为一个阶跃函数,当 x 为负时返回 0,当 x 为正时返回 1

import jax
import numpy as np
import jax.numpy as jnp

def f(x):
  return (x > 0).astype(float)

df = jax.vmap(jax.grad(f))

x = jnp.array([-1.0, -0.5, 0.0, 0.5, 1.0])

print(f"f(x)  = {f(x)}")
# f(x)  = [0. 0. 0. 1. 1.]

print(f"df(x) = {df(x)}")
# df(x) = [0. 0. 0. 0. 0.]

梯度到处为零的事实乍一看可能令人困惑:毕竟输出确实随着输入的变化而变化,梯度怎么可能是零呢?然而,事实证明在这种情况下零就是正确的结果。

原因是什么?记住,微分测量的是 x 发生无穷小变化时 f 的变化量。对于 x=1.0f 返回 1.0。如果我们微调 x 使其略大或略小,输出都不会改变,因此根据定义,grad(f)(1.0) 应为零。这一逻辑对所有大于零的 x 值均成立。类似地,对于所有小于零的 x 值,输出均为零,微调 x 不会改变输出,因此梯度为零。这就剩下棘手的 x=0 情况。当然,如果你向上微调 x,输出会发生改变,但这很成问题:x 的无穷小变化产生函数值的有限变化,这意味着梯度未定义。幸运的是,我们还有另一种测量梯度的方法:我们将函数向下微调,此时输出不变,因此梯度为零。JAX 和其他自动微分系统倾向于以这种方式处理不连续性:如果正梯度和负梯度不一致,但其中一个有定义而另一个没有,我们使用有定义的那个。根据此定义,该函数的梯度在数学和数值上均处处为零。

问题的根源在于函数在 x = 0 处有一个不连续点。这里的 f 本质上是一个 Heaviside 阶跃函数,我们可以使用 Sigmoid 函数 作为平滑替代方案。当 x 远离零时,Sigmoid 近似等于阶跃函数,但它将 x = 0 处的不连续性替换为平滑、可微分的曲线。结果,通过使用 jax.nn.sigmoid(),我们得到了具有定义明确梯度的类似计算。

def g(x):
  return jax.nn.sigmoid(x)

dg = jax.vmap(jax.grad(g))

x = jnp.array([-10.0, -1.0, 0.0, 1.0, 10.0])

with np.printoptions(suppress=True, precision=2):
  print(f"g(x)  = {g(x)}")
  # g(x)  = [0.   0.27 0.5  0.73 1.  ]

  print(f"dg(x) = {dg(x)}")
  # dg(x) = [0.   0.2  0.25 0.2  0.  ]

jax.nn 子模块还有其他常用排序相关函数的平滑版本,例如 jax.nn.softmax() 可以替代 jax.numpy.argmax()jax.nn.soft_sign() 可以替代 jax.numpy.sign()jax.nn.softplus()jax.nn.squareplus() 可以替代 jax.nn.relu() 等。

如何将 JAX 追踪器 (Tracer) 转换为 NumPy 数组?#

在运行时检查转换后的 JAX 函数时,你会发现数组值被 jax.core.Tracer 对象替换了。

@jax.jit
def f(x):
  print(type(x))
  return x

f(jnp.arange(5))

这会打印以下内容:

<class 'jax.interpreters.partial_eval.DynamicJaxprTracer'>

常见问题是此类追踪器如何转回常规 NumPy 数组。简而言之,追踪器不可能转换为 NumPy 数组,因为追踪器是对给定形状和数据类型的所有 可能值 的抽象表示,而 NumPy 数组是该抽象类的一个具体成员。有关 JAX 转换上下文中追踪器工作方式的更多讨论,请参阅 JIT 机制

将追踪器转换回数组的问题通常出现在另一个目标的上下文中,即在运行时访问计算中的中间值。例如:

有关运行时回调及其使用的更多信息和示例,请参阅 JAX 中的外部回调

为什么某些 CUDA 库无法加载/初始化?#

在解析动态链接库时,JAX 使用通常的 动态链接器搜索模式。JAX 设置 RPATH 指向 pip 安装的 NVIDIA CUDA 包的 JAX 相关位置,并优先使用它们。如果 ld.so 在其通常的搜索路径中无法找到 CUDA 运行时库,则必须在 LD_LIBRARY_PATH 中明确包含这些库的路径。确保 CUDA 文件可被发现的最简单方法是直接安装 nvidia-*-cu12 pip 包,这些包包含在标准的 jax[cuda_12] 安装选项中。

偶尔,即使你确保了运行时库可被发现,在加载或初始化时仍可能出现问题。此类问题的常见原因是运行时 CUDA 库初始化时内存不足。有时这是因为 JAX 为了提高执行速度,会预分配过大比例的当前可用设备内存,导致运行时 CUDA 库初始化时剩余内存不足。

在使用多个 JAX 实例、将 JAX 与执行自身预分配的 TensorFlow 并行使用,或者在 GPU 被其他进程大量占用的系统上运行 JAX 时,这种情况尤为可能。如有疑问,请尝试通过减少预分配来重新运行程序,方法是将 XLA_PYTHON_CLIENT_MEM_FRACTION 从默认的 .75 调低,或者设置 XLA_PYTHON_CLIENT_PREALLOCATE=false。有关详细信息,请参阅 JAX GPU 内存分配 页面。

基准测试 JAX 代码#

已移至 基准测试 JAX 代码

缓冲区捐赠 (Buffer donation)#

已移至 缓冲区捐赠