JAX 调试标志#
JAX 提供了一些标志和上下文管理器,使错误捕获变得更加容易。
jax_debug_nans 配置选项和上下文管理器#
摘要:启用 jax_debug_nans 标志,以便在 jax.jit 编译的代码中自动检测 NaN 的产生。
jax_debug_nans 是一个 JAX 标志,启用后,当计算产生 NaN 时,它会导致程序立即报错。开启此选项会为 XLA 产生的每个浮点类型值增加 NaN 检查。这意味着在未受 @jax.jit 装饰的代码中,每个原始操作产生的值都会被拉回到主机,并作为 ndarray 进行检查。
对于 @jax.jit 装饰下的代码,每个 @jax.jit 函数的输出都会被检查;如果存在 NaN,它将以非优化(de-optimized)的逐操作(op-by-op)模式重新运行该函数,这实际上相当于每次移除一层 @jax.jit。
可能会出现一些棘手的情况,例如 NaN 仅在 @jax.jit 下才会出现,而在非优化模式下不会产生。在这种情况下,你会看到警告信息输出,但代码将继续执行。
如果 NaN 是在梯度评估的反向传播过程中产生的,当异常在堆栈跟踪上方几层被抛出时,你将处于 backward_pass 函数中,这本质上是一个简单的 jaxpr 解释器,它会逆序遍历原始操作序列。
用法#
如果你想追踪 NaN 在函数或梯度中的产生位置,可以通过以下任一方式开启 NaN 检查器:
在
jax.debug_nans上下文管理器中运行代码,使用with jax.debug_nans(True):;设置
JAX_DEBUG_NANS=True环境变量;在主文件顶部附近添加
jax.config.update("jax_debug_nans", True);在主文件中添加
jax.config.parse_flags_with_absl(),然后使用命令行标志(如--jax_debug_nans=True)设置该选项;
示例#
import jax
import jax.numpy as jnp
import traceback
jax.config.update("jax_debug_nans", True)
def f(x):
w = 3 * jnp.square(x)
return jnp.log(-w)
# The stack trace is very long so only print a couple lines.
try:
f(5.)
except FloatingPointError as e:
print(traceback.format_exc(limit=2))
Invalid nan value encountered in the output of a jax.jit function. Calling the de-optimized version.
Traceback (most recent call last):
File "/tmp/ipykernel_1736/1479925735.py", line 12, in <module>
f(5.)
File "/tmp/ipykernel_1736/1479925735.py", line 8, in f
return jnp.log(-w)
^^^^^^^^^^^
FloatingPointError: invalid value (nan) encountered in log
生成的 NaN 已被捕获。通过运行 %debug,我们可以获得事后调试器(post-mortem debugger)。如下例所示,这也适用于 @jax.jit 装饰的函数。
jax.jit(f)(5.)
Invalid nan value encountered in the output of a jax.jit function. Calling the de-optimized version.
Invalid nan value encountered in the output of a jax.jit function. Calling the de-optimized version.
---------------------------------------------------------------------------
FloatingPointError Traceback (most recent call last)
Cell In[2], line 1
----> 1 jax.jit(f)(5.)
[... skipping hidden 5 frame]
Cell In[1], line 8, in f(x)
6 def f(x):
7 w = 3 * jnp.square(x)
----> 8 return jnp.log(-w)
[... skipping hidden 5 frame]
File ~/checkouts/readthedocs.org/user_builds/jax/envs/latest/lib/python3.12/site-packages/jax/_src/numpy/ufuncs.py:491, in log(x)
456 @export
457 @jit(inline=True)
458 def log(x: ArrayLike, /) -> Array:
459 """Calculate element-wise natural logarithm of the input.
460
461 JAX implementation of :obj:`numpy.log`.
(...) 489 Array(True, dtype=bool)
490 """
--> 491 out = lax.log(*promote_args_inexact('log', x))
492 jnp_error._set_error_if_nan(out)
493 return out
[... skipping hidden 7 frame]
File ~/checkouts/readthedocs.org/user_builds/jax/envs/latest/lib/python3.12/site-packages/jax/_src/pjit.py:171, in _run_python_pjit(p, args_flat, fun, args, kwargs)
169 except api_util.InternalFloatingPointError as e:
170 if getattr(fun, '_apply_primitive', False):
--> 171 raise FloatingPointError(
172 f"invalid value ({e.ty}) encountered in {fun.__qualname__}") from None
173 api_util.maybe_recursive_nan_check(e, fun, args, kwargs) # should always raise.
174 raise RuntimeError("Internal error") from e # fall-back error to be safe.
FloatingPointError: invalid value (nan) encountered in log
当此代码在 @jax.jit 函数的输出中看到 NaN 时,它会调用非优化代码,因此我们仍然可以获得清晰的堆栈跟踪。我们可以使用 %debug 运行事后调试器来检查所有值,以找出错误所在。
jax.debug_nans 上下文管理器可用于激活/停用 NaN 调试。由于我们在上面通过 jax.config.update 激活了它,现在让我们停用它。
with jax.debug_nans(False):
print(jax.jit(f)(5.))
nan
jax_debug_nans 的优势与局限性#
优点#
易于应用
能够精确检测 NaN 的产生位置
抛出标准 Python 异常,并与 PDB 事后调试兼容
局限性#
急切(eager)地重新运行函数可能会很慢。在不调试时,你不应该开启 NaN 检查器,因为它会引入大量的设备与主机间通信,并导致性能下降。
对误报(例如有意创建的 NaN)报错
jax_debug_infs 配置选项和上下文管理器#
jax_debug_infs 的工作方式与 jax_debug_nans 类似。jax_debug_infs 通常需要与 jax_disable_jit 结合使用,因为 Inf 可能不会像 NaN 那样级联传播到输出。或者,可以使用 jax.experimental.checkify 来查找中间结果中的 Inf。
jax_debug_infs 的完整文档即将发布。
jax_disable_jit 配置选项和上下文管理器#
摘要:启用 jax_disable_jit 标志以禁用 JIT 编译,从而允许使用传统的 Python 调试工具,如 print 和 pdb。
jax_disable_jit 是一个 JAX 标志,启用后,会在整个 JAX 中禁用 JIT 编译(包括在控制流函数中,如 jax.lax.cond 和 jax.lax.scan)。
用法#
你可以通过以下方式禁用 JIT 编译:
设置
JAX_DISABLE_JIT=True环境变量;在主文件顶部附近添加
jax.config.update("jax_disable_jit", True);在主文件中添加
jax.config.parse_flags_with_absl(),然后使用命令行标志(如--jax_disable_jit=True)设置该选项;
示例#
import jax
jax.config.update("jax_disable_jit", True)
def f(x):
y = jnp.log(x)
if jnp.isnan(y):
breakpoint()
return y
jax.jit(f)(-2.) # ==> Enters PDB breakpoint!
jax_disable_jit 的优势与局限性#
优势#
易于应用
允许使用 Python 内置的
breakpoint和print抛出标准 Python 异常,并与 PDB 事后调试兼容
局限性#
在没有 JIT 编译的情况下运行函数可能会很慢