常见问题解答 (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 识别出 log 是 exp 的逆运算,从而在编译后的函数中移除了这些操作,直接返回输入值。在这种情况下,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)
延伸阅读
为什么基于排序顺序的函数的梯度为零?#
如果你定义的函数处理输入的操作依赖于输入的相对顺序(例如 max、greater、argsort 等),你可能会惊讶地发现梯度到处都是零。这是一个例子,我们将 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.0,f 返回 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.debug.print()。如果你希望在转换后的 JAX 函数中调用非 JAX 代码,可以考虑使用
jax.pure_callback(),示例请参见 纯回调示例。如果你希望在运行时输入或输出数组缓冲区(例如从文件加载数据,或将数组内容记录到磁盘),可以考虑使用
jax.experimental.io_callback(),示例请参见 IO 回调示例。
有关运行时回调及其使用的更多信息和示例,请参阅 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)#
已移至 缓冲区捐赠。