首页
/ JAX 编译函数调试实战:jax.debug.print 与 jax.debug.breakpoint 完全指南

JAX 编译函数调试实战:jax.debug.print 与 jax.debug.breakpoint 完全指南

2026-09-09 12:04:05作者:董宙帆

导读

在 JAX 中,jax.jitjax.pmapjax.pjit 等变换会延迟数值求值,导致 Python 内建的 printpdb 在编译函数内部失效。本文基于 docs/debugging/print_breakpoint.md 展开,系统讲解 jax.debug.printjax.debug.breakpoint 的用法、底层原理、排序与副作用等"尖锐边角"(sharp bits)及性能影响,并结合 jax/debug.pyjax/_src/debugging.pyjax/_src/debugger/core.py 等源码给出实现级证据。读完本文,你将掌握在 JIT 编译函数中打印张量中间值、交互式暂停检查调用栈、以及在 grad/vmap/pmap/pjit 等变换下的正确调试姿势。


1. 为什么编译函数里需要专用的调试工具

JAX 的核心设计是"组合式变换":jax.jit 会把你的 Python 函数追踪(trace)成中间表示(jaxpr),再交给 XLA 编译成可执行程序。数值计算被延迟到设备(GPU/TPU)上异步执行,Python 侧的 print 在追踪阶段只会打印"追踪器"(tracer),而不是真实数值;jax.pmap 等并行变换还会把同一个函数分发到多设备上执行。因此传统 Python 调试手段在编译函数中失效。

jax.debug 包(对应源码 jax/debug.py)正是为此设计的一组工具,其中:

  • jax.debug.print:把被追踪的数组值打印到 stdout;
  • jax.debug.breakpoint:暂停编译函数的执行,进入交互式调试器检查调用栈中的值;
  • 此外还有 jax.debug.callback(底层回调)、jax.debug.logjax.debug.inspect_array_shardingjax.debug.visualize_array_sharding 等配套工具。

jax/_src/debugging.py 可以看到,调试回调在 JAX 的效果系统(effects system)中被建模为 DebugEffectOrderedDebugEffect 两类效果,并显式声明它们允许出现在控制流、重计算(remat)、自定义导数等场景中。这从底层解释了为什么 jax.debug.print 能穿透 jitvmapgrad 等变换。


2. 用 jax.debug.print 打印追踪值

2.1 最小示例

jax.jitjax.pmap 装饰的函数中打印被追踪的数组值:

import jax
import jax.numpy as jnp

@jax.jit
def f(x):
  jax.debug.print("🤯 {x} 🤯", x=x)
  y = jnp.sin(x)
  jax.debug.print("🤯 {y} 🤯", y=y)
  return y

f(2.)
# Prints:
# 🤯 2.0 🤯
# 🤯 0.9092974662780762 🤯

jax.debug.print 语义上等价于下面的 Python 函数,区别在于它可以被 JAX 编排(staged out)并随变换一起执行:

def debug.print(fmt: str, *args: PyTree[Array], **kwargs: PyTree[Array]) -> None:
  print(fmt.format(*args, **kwargs))

关键约束:fmt 不能是 f-string。 因为 f-string 在调用瞬间就完成格式化,而 jax.debug.print 需要把格式化推迟到数值真正求值之后。正确写法是把占位符留到格式串里、把变量作为关键字参数传入:jax.debug.print("hello {bar}", bar=bar)。这一点在源码 jax/_src/debugging.py 中也有体现:_DebugPrintFormatChecker 会校验未使用的参数并抛出 ValueError,提示"你可能是把 f-string 传给了 jax.debug.print"。

从实现细节看(jax/_src/debugging.py),debug_print 支持两种调用形式:

  • 单次调用:jax.debug.print("hello {x}", x=42, ordered=True)(将 ordered/partitioned 选项与打印参数混用属于软弃用);
  • 两次调用(推荐):jax.debug.print(ordered=True)("hello {x}", x=42),选项放在第一次调用,格式串与参数放在第二次。

2.2 什么时候用 jax.debug.print

  • 动态(被追踪)的数组值:在 jitvmappmappjit 等变换内部,必须用 jax.debug.print
  • 静态值(数组形状、dtype 等):这些在追踪期就已知,直接用普通 Python print 即可。

对于 jax.gradjax.vmap 这类变换,Python 内建 print 有时也能打印出数值,但在 jax.jit/jax.pmap 下会失效——统一使用 jax.debug.print 可以保证行为一致。

2.3 为什么用 jax.debug.print:揭示计算求值方式

jax.debug.print 能暴露"计算实际上是怎样被求值的",例如对比 jax.vmapjax.lax.map 的求值顺序:

xs = jnp.arange(3.)

def f(x):
  jax.debug.print("x: {}", x)
  y = jnp.sin(x)
  jax.debug.print("y: {}", y)
  return y

jax.vmap(f)(xs)
# Prints: x: 0.0
#         x: 1.0
#         x: 2.0
#         y: 0.0
#         y: 0.841471
#         y: 0.9092974

jax.lax.map(f, xs)
# Prints: x: 0.0
#         y: 0.0
#         x: 1.0
#         y: 0.841471
#         x: 2.0
#         y: 0.9092974

注意两次打印顺序不同:vmap 下先打印所有 x 再打印所有 y,而 lax.map 下是交替打印。这揭示了一个重要事实——jax.debug.print 的输出不遵守 JAX 通常的语义保证(例如 jax.vmap(f)(xs)jax.lax.map(f, xs) 计算结果相同但求值方式不同)。正是这些求值顺序细节,恰恰是调试时想看到的东西。

结论:jax.debug.print 只用于调试,不要在语义保证重要的场景中使用。

从源码看(jax/_src/debugging.py),vmap 对调试回调的 batching 规则会**沿映射轴展开(unroll)**回调:debug_batching_rule 逐索引取轴切片并分别绑定原语,这正是上面"先全部 x 后全部 y"顺序的由来。而 lax.map 是顺序循环,所以 x、y 交替出现。

2.4 更多 jax.debug.print 示例

jax.pmap 下打印

jax.pmap 并行执行时,打印可能乱序

xs = jnp.arange(2.)

def f(x):
  jax.debug.print("x: {}", x)
  return x

jax.pmap(f)(xs)
# Prints: x: 0.0
#         x: 1.0
# OR
# Prints: x: 1.0
#         x: 0.0

jax.grad 下打印

jax.grad 下,jax.debug.print 只在正向传播(forward pass)打印

def f(x):
  jax.debug.print("x: {}", x)
  return x * 2.

jax.grad(f)(1.)
# Prints: x: 1.0

这与 Python 内建 printgrad 下的行为类似,但用 jax.debug.print 后,即使外层再套 jax.jit,行为也保持一致。

若想在**反向传播(backward pass)**打印梯度,可以用 jax.custom_vjp 在反向函数里插入打印:

@jax.custom_vjp
def print_grad(x):
  return x

def print_grad_fwd(x):
  return x, None

def print_grad_bwd(_, x_grad):
  jax.debug.print("x_grad: {}", x_grad)
  return (x_grad,)

print_grad.defvjp(print_grad_fwd, print_grad_bwd)


def f(x):
  x = print_grad(x)
  return x * 2.

jax.grad(f)(1.)
# Prints: x_grad: 2.0

在其他变换下打印

jax.debug.print 同样适用于 pjit 等其它变换。

2.5 更灵活的控制:jax.debug.callback

事实上,jax.debug.printjax.debug.callback 的一个薄封装。callback 允许你传入任意 Python 可调用对象,对格式化和输出类型有更强的控制:

def callback(fun: Callable, *args: PyTree[Array], **kwargs: PyTree[Array]) -> None:
  fun(*args, **kwargs)
  return None

jax/_src/debugging.py 的签名看,jax.debug.callback(callback, *args, ordered=False, partitioned=False, **kwargs)partitioned=True 时只打印本设备的本地分片(避免对操作数的 all-gather);False 时打印逻辑操作数(需先 all-gather)。回调同样支持单次调用与两次调用两种形式。

注意:回调只应用于无害的调试输出(打印、绘图等)。不要用它做计时等操作——回调可能被重排且是异步的(见下文"尖锐边角")。jax.experimental.io_callbackjax.pure_callback 是面向有副作用/纯函数的更正式回调方案,可参考 jax/debug.py 中的相关文档说明。


3. jax.debug.print 的"尖锐边角"(Sharp bits)

像大多数 JAX API 一样,jax.debug.print 用不好也会"割到手"。以下是官方文档与源码共同确认的注意事项。

3.1 打印结果可能乱序

当多次 jax.debug.print 调用涉及互不依赖的参数时,它们可能在编排(staged out)时被编译器重排,例如在 jax.jit 下:

@jax.jit
def f(x, y):
  jax.debug.print("x: {}", x)
  jax.debug.print("y: {}", y)
  return x + y

f(2., 3.)
# Prints: x: 2.0
#         y: 3.0
# OR
# Prints: y: 3.0
#         x: 2.0

原因:编译器拿到的是被编排计算的功能化表示(functional representation),Python 函数的命令式顺序被丢弃,只保留数据依赖关系。对纯函数用户不可见,但涉及打印这类副作用时就能观察到差异。

解决办法:使用 jax.debug.print(..., ordered=True) 保证打印相对顺序与源码一致。但 ordered=Truejax.pmap 及其他涉及并行的变换下会报错,因为并行执行下无法保证顺序。从源码看,ordered=True 时使用 OrderedDebugEffectjax/_src/debugging.py),并走带 token 的 emit_python_callback 路径(jax/_src/debugging.py),token 机制保证了有序回调的相对顺序;同时 debug_callbackorderedtoken 参数互斥(jax/_src/debugger/core.py)。

3.2 计算扰动(Computation perturbation)

加入 jax.debug.printjax.debug.breakpoint 会改变交给 XLA 编译的计算本身。由于 XLA 编译期可能进行不同的算子融合(fusion),数值可能与无调试语句的版本出现细微差异。排查数值问题时请牢记:加调试语句这一行为本身可能影响你正在调查的行为。

3.3 异步回调

取决于后端,jax.debug.print 可能异步发生(不在主线程),即 JAX 函数返回后打印才出现在屏幕上:

@jax.jit
def f(x):
  jax.debug.print("x: {}", x)
  return x

f(2.).block_until_ready()
# <do something else>
# Prints: x: 2.

要阻塞等待函数内的打印完成,可调用 jax.effects_barrier()(对应实现见 jax/_src/api.py),它会等待函数中剩余的副作用全部完成:

@jax.jit
def f(x):
  jax.debug.print("x: {}", x)
  return x

f(2.).block_until_ready()
jax.effects_barrier()
# Prints: x: 2.
# <do something else>

3.4 性能影响

不必要的物化(Unnecessary materialization)

jax.debug.print 设计上追求最小的性能足迹,但可能干扰编译器优化、影响内存画像:

def f(w, b, x):
  logits = w.dot(x) + b
  jax.debug.print("logits: {}", logits)
  return jax.nn.relu(logits)

上例在线性层与激活函数之间打印中间值。XLA 本可通过融合优化避免把 logits 物化到内存,但打印强制物化了这些中间值,可能拖慢程序、增加内存占用。此外,在 jax.pjit 下使用 jax.debug.print 还会发生一次全局同步,把值物化到单个设备上。

回调开销(Callback overhead)

jax.debug.print 本质上需要在加速器与主机(host)之间通信:无论 GPU 还是 TPU,底层机制不同,但都需要把打印值从设备拷贝回主机。CPU 场景下该开销较小。jax.pjit 下的全局同步也会带来额外开销。

从源码看,调试回调的实现(jax/_src/debugging.py)在 CPU/GPU 上注册了可缓存的 lowering,而 TPU 因使用 channel ID 不可缓存cacheable=False,见 jax/_src/debugging.py);同时调试原语的抽象求值声明了效果(effect),避免被死代码消除(DCE)掉。

3.5 优劣势小结

优势

  • 打印调试简单直观;
  • jax.debug.callback 可用于其它无害副作用。

局限

  • 添加打印语句是手工过程;
  • 可能存在性能影响。

4. 用 jax.debug.breakpoint() 交互式检查值

4.1 基本用法

jax.debug.breakpoint()暂停 JAX 程序的执行,让你检查调用栈中的值:

@jax.jit
def f(x):
  y, z = jnp.sin(x), jnp.cos(x)
  jax.debug.breakpoint()
  return y * z

f(2.) # ==> Pauses during execution!

JAX 交互式调试器暂停在断点处

jax.debug.breakpoint() 本质上是 jax.debug.callback(...) 的一个应用,额外捕获了调用栈信息,因此继承了 jax.debug.print 的所有变换行为(例如 vmap 会沿映射轴展开断点)。

从源码看(jax/_src/debugger/core.py),breakpoint 支持以下关键字参数:

  • backend:选择调试器后端,默认选取优先级最高的已注册调试器(如 CLI 调试器);
  • filter_frames:是否过滤掉 JAX 内部栈帧(默认 True),可能影响 Flax 等库的栈帧过滤;
  • num_frames:可供交互检查的栈帧数量上限;
  • ordered:是否保证该断点相对其他有序 breakpoint/print 调用的顺序;
  • token:与 ordered 互斥的替代方案——传入一个 JAX 数组(或 pytree),断点会在该值计算完成后运行,并把 token 原样返回(若返回值未被后续计算使用,整个计算会被剪枝,断点不运行)。

4.2 调试器命令

命中断点后,你会看到一个类似 pdb 的提示符(CLI 调试器的提示符为 (jdb),见 jax/_src/debugger/cli_debugger.py)。与 pdb 不同,你不能单步执行,但可以恢复执行。可用命令如下:

命令 功能
help 打印可用命令
p 计算表达式并打印结果
pp 计算表达式并漂亮打印结果
u(p) 向上移动栈帧
d(own) 向下移动栈帧
w(here) / bt 打印回溯(backtrace)
l(ist) 打印代码上下文
c(ont(inue)) 恢复程序执行
q(uit) / exit 退出程序(TPU 上不可用

在 CLI 调试器实现中(jax/_src/debugger/cli_debugger.py),p 类求值会把当前帧的 globalslocals 合并后 evalbt 会逆序遍历保存的帧打印回溯;l 会打印当前帧附近 ±2 行的源码。此外调试器还注册了 Colab 与 Web 后端(见 jax/_src/debugger/init.py),可适配不同交互环境。

4.3 与 jax.lax.cond 配合检测 NaN/Inf

把断点与 jax.lax.cond 结合,可以方便地检测 naninf

def breakpoint_if_nonfinite(x):
  is_finite = jnp.isfinite(x).all()
  def true_fn(x):
    pass
  def false_fn(x):
    jax.debug.breakpoint()
  lax.cond(is_finite, true_fn, false_fn, x)

@jax.jit
def f(x, y):
  z = x / y
  breakpoint_if_nonfinite(z)
  return z

f(2., 0.) # ==> Pauses during execution!

x / y 产生非有限值时,false_fn 触发断点,让你在现场检查调用栈中的中间值。

4.4 断点的"尖锐边角"

由于 jax.debug.breakpoint 只是 jax.debug.callback 的应用,它继承 jax.debug.print 的全部"尖锐边角"(见上文第 3 节),并额外有两处注意:

  • 物化更多中间值:断点会强制物化调用栈中的所有值,物化量比 jax.debug.print 更大;
  • 运行时开销更高:断点可能需要把 JAX 程序中所有中间值从设备拷贝到主机。

4.5 优劣势小结

优势

  • 简单直观且(某种程度上)标准化;
  • 可以同时检查调用栈上下的大量值。

局限

  • 定位错误源头可能需要设置多个断点;
  • 物化大量中间值,开销较大。

5. 实战要点速查

  1. 编译函数内打印数值:一律用 jax.debug.print,格式串用普通字符串 + {} 占位符,变量以关键字传入;禁用 f-string
  2. 打印静态信息(shape、dtype):普通 Python print 即可。
  3. 保持打印顺序:使用 ordered=True;但并行变换(pmap 等)下会报错,属预期行为。
  4. 等待异步打印完成:在函数返回后调用 jax.effects_barrier()
  5. 交互式检查中间值:用 jax.debug.breakpoint(),并通过 p/pp/u/d/bt/l/c/q 等命令导航;与 jax.lax.cond 结合可自动捕获 NaN/Inf 现场。
  6. 反向传播打印梯度:通过 jax.custom_vjp 在反向函数中插入 jax.debug.print
  7. 记住调试本身的扰动:加调试语句可能改变 XLA 融合、影响数值与性能;正式运行前记得移除调试代码。
  8. 查看 shardingjax.debug.inspect_array_sharding / jax.debug.visualize_array_sharding 可在 jit 内部检查中间值的分片信息(jax/_src/debugging.py),配合 AOT API(jit(f).lower(...))可提前触发回调。

6. 相关资源

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
docsdocs
暂无描述
Markdown
900
5.83 K
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.35 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
927
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.89 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
602
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
396
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.04 K
526