首页
/ JAX 运行时调试实战指南:深入掌握 `jax.debug` 模块的打印、断点与分片检查

JAX 运行时调试实战指南:深入掌握 `jax.debug` 模块的打印、断点与分片检查

2026-09-09 17:22:41作者:盛欣凯Ernestine

jax.debug 是 JAX 提供的一组运行时调试工具,专门用于在 jax.jitjax.pmapjax.pjit 等延迟执行(staged)的变换中打印数值、暂停执行并检查调用栈中的中间值,以及检视数组的分片(sharding)信息。读完本文,你将掌握 jax.debug.printjax.debug.callbackjax.debug.breakpoint 三大运行时值调试 API 与 inspect_array_shardingvisualize_array_shardingvisualize_sharding 三大分片调试 API 的完整用法、边界行为与底层实现原理,能够在分布式与编译场景下高效定位数值异常(如 NaN/Inf)与分片问题。

jax.debug 模块总览

jax.debug 是 JAX 对外公开的调试 API 入口,其导出定义位于 jax/debug.py

__all__ = ["callback", "print", "log", "DebugEffect", "OrderedDebugEffect",
           "visualize_array_sharding",
           "inspect_array_sharding", "visualize_sharding", "breakpoint"]

从导出列表可以看出,模块按功能可分为两组(对应 docs/jax.debug.rst 的组织结构):

  • 运行时值调试(Runtime value debugging)callbackprintbreakpoint,以及内部使用的效果标记 DebugEffectOrderedDebugEffectlog。官方完整指南见 docs/debugging/print_breakpoint.md
  • 分片调试(Sharding debugging)inspect_array_shardingvisualize_array_shardingvisualize_sharding,用于在(以及不在)staged 函数内部检视和可视化数组分片。

全部实现位于 jax/_src/debugging.py(打印、回调与分片可视化)与 jax/_src/debugger/(断点调试器)中,jax/debug.py 只是薄薄的转发层。


一、jax.debug.print:在编译函数内打印数值

1.1 为什么普通 printjax.jit 中失效

JAX 的变换分为两类:jax.gradjax.vmap 会保留 Python 层的即时执行语义,此时内置 print 能打印数值;而 jax.jitjax.pmap 会延迟数值求值(先 trace 成计算图,再交给 XLA 编译执行),函数体内的 Python 代码只在追踪(tracing)阶段运行一次,内置 print 打印的是 tracer 而非真实数值。因此在这类编译场景下必须使用 jax.debug.print

1.2 基本用法

jax.debug.print 的语义上等价于 print(fmt.format(*args, **kwargs)),区别在于它可以被 staged 并被 JAX 变换

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 🤯

重要约束:fmt 不能是 f-string。f-string 会在 Python 层立即完成格式化,而 jax.debug.print 需要把格式化延迟到运行时(设备端数值就绪之后)。应写成 jax.debug.print("hello {bar}", bar=bar) 而非 jax.debug.print(f"hello {bar}")。从源码看,_DebugPrintFormatCheckerjax/_src/debugging.py)会主动检查未使用的参数并抛出错误提示 "You may be passing an f-string (i.e, f\"{x}\")...",帮助你尽早发现误用。

1.3 何时使用 jax.debug.print

  • 需要打印动态(被 trace 的)数组值时,在 jitvmap 等变换内使用 jax.debug.print
  • 只打印静态信息(如 shape、dtype)时,普通 Python print 即可,无需引入调试 API。

1.4 揭示计算求值顺序

jax.debug.print 的独特价值在于能暴露"计算到底如何被求值"。以下面代码为例(同样出自 docs/debugging/print_breakpoint.md):

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 则是每个元素"x、y"交替(循环语义)。这意味着 jax.debug.print 的输出不遵循 JAX 通常的语义保证(例如 vmap(f)(xs)lax.map(f, xs) 计算结果相同但求值方式不同),而这种求值顺序细节恰恰是调试时最想看到的。因此,jax.debug.print 只用于调试,不要依赖它的输出顺序来保证任何语义。

1.5 各变换下的行为

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 下只在正向传播打印

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 后行为依然不变。若要在反向传播打印梯度,借助 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 等其它变换。

1.6 两种调用形式与参数

jax.debug.print 支持两种调用形式(源码 jax/_src/debugging.py):

# 1. 两段式调用(推荐):选项放在第一次调用,fmt 与参数放在第二次
jax.debug.print(ordered=True)("hello {x}", x=42)

# 2. 单段式调用(软弃用):选项混入 fmt 的 kwargs
jax.debug.print("hello {x}", x=42, ordered=True)

关键参数:

参数 含义
fmt 格式字符串,语义同 str.format,支持 {} 位置占位与 {name} 命名占位
*args / **kwargs 待格式化的参数,会作为 PyTree 扁平化后传入回调
ordered 是否保证该打印与其它 ordered=True 的打印/断点之间的相对顺序(见 1.8 节)
partitioned True 时只打印本地分片(避免 all-gather);为 False 时打印逻辑整体值(需要先 all-gather 操作数)
skip_format_check True 时跳过格式串校验,适用于 Pallas TPU kernel 等场景(标量参数会打印在格式串之后)

1.7 jax.debug.callback:更底层的通用回调

jax.debug.print 本质上是 jax.debug.callback 的便捷封装。jax.debug.callback(fun, *args, **kwargs) 语义上等价于 fun(*args, **kwargs); return None,但可以被 staged 出来参与 JAX 变换。它给你对格式化方式和输出目标的完全控制:

import jax

def callback(fun, *args, **kwargs):
  # 在 staged 程序中以纯函数语义调用 fun
  fun(*args, **kwargs)

jax.debug.callback 同样支持两种调用形式(两段式推荐),参数 orderedpartitioned 含义与 print 一致(jax/_src/debugging.py)。它只适合无副作用的调试输出(打印、绘图等);不要用它做计时等操作——因为回调可能被重排且异步执行。

1.8 Sharp bits:jax.debug.print 的锋利边缘

打印结果可能被重排

当多个 jax.debug.print 的参数之间没有数据依赖时,staged 之后编译器拿到的是一份函数式表示,Python 命令式顺序已丢失,只剩数据依赖关系,因此打印可能乱序:

@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   或   y: 3.0 / x: 2.0

如需保持代码中书写顺序,使用 jax.debug.print(..., ordered=True);但 ordered=Truejax.pmap 等并行变换下会报错,因为并行执行无法保证顺序。从源码看,ordered 对应 OrderedDebugEffect 效果类型,在 MLIR lowering 时会通过 token 机制串行化(jax/_src/debugging.py)。

计算扰动(Computation perturbation)

加入 jax.debug.printjax.debug.breakpoint 会改变 XLA 实际编译的计算内容,XLA 可能做出不同的算子融合决策,从而与无调试语句的代码产生数值差异。排查数值问题时需警惕:加调试语句这个动作本身就可能改变你要调查的行为。

异步回调

取决于后端,jax.debug.print 可能异步执行(不在主线程),函数已返回后数值才打印到屏幕:

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

f(2.).block_until_ready()
jax.effects_barrier()
# Prints: x: 2.

性能影响

  • 不必要的物化(materialization):打印中间值会强制编译器物化该中间结果。例如 logits = w.dot(x) + b; jax.debug.print("logits: {}", logits); return jax.nn.relu(logits),XLA 本可融合优化、避免将 logits 写入内存,而打印会迫使它被物化,可能拖慢程序、增加内存占用。在 jax.pjit 下还会发生全局同步,把值物化到单个设备上。
  • 回调开销jax.debug.print 必然引入加速器与主机之间的通信(将设备值拷贝回主机),GPU/TPU 与 CPU 的底层机制不同,CPU 开销较小;jax.pjit 场景下全局同步还会叠加额外开销。

优势与局限

优势:打印式调试简单直观;jax.debug.callback 可承载其它无害副作用。 局限:需要手动插入打印语句;存在性能影响。


二、jax.debug.breakpoint:交互式检查调用栈

2.1 基本用法

jax.debug.breakpoint() 会在编译函数执行时暂停程序,呈现一个类似 pdb 的交互提示符,允许你检查调用栈中的值:

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

f(2.) # ==> 执行到断点时暂停!

jax.debug.breakpoint 触发后进入的交互式调试会话

pdb 不同,你不能单步执行,但可以恢复执行。从实现看,jax.debug.breakpoint() 只是 jax.debug.callback(...) 的一个应用——回调内容为捕获调用栈信息(jax/_src/debugger/core.py),因此它天然继承 jax.debug.print 的全部变换行为,例如对 jax.debug.breakpoint()vmap 会在映射轴上展开。

2.2 调试器命令

进入调试器后支持以下命令(对应 pdb 风格,提示符为 (jdb),见 jax/_src/debugger/cli_debugger.py):

命令 作用
help 打印可用命令
p 求值表达式并打印结果
pp 求值表达式并 pretty-print 结果
u(p) 上移一个栈帧
d(own) 下移一个栈帧
w(here) / bt 打印回溯(backtrace)
l(ist) 打印代码上下文
c(ont(inue)) 恢复程序执行
q(uit) / exit 退出程序(TPU 上不可用)

默认情况下调试器会过滤 JAX 内部栈帧(filter_frames=True),让你聚焦用户代码;num_frames 可限制可检查的栈帧数量。底层调试器通过注册表机制选择(jax/_src/debugger/core.py),CLI 调试器为默认后端,Colab 与 Web 调试器也在同一目录下实现。

2.3 实战示例:配合 jax.lax.cond 检测 NaN/Inf

将断点与 jax.lax.cond 组合,可以精确地在数值非有限时暂停:

import jax
from jax import lax
import jax.numpy as jnp

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.) # ==> 除零产生 inf,执行到断点时暂停!

2.4 Sharp bits 与权衡

由于 jax.debug.breakpoint 只是 jax.debug.callback 的应用,它继承了 jax.debug.print 的所有 sharp bits,并额外有两点:

  • 物化更多的中间值:因为它强制把调用栈中的所有值都物化出来;
  • 运行时开销更大:需要把 JAX 程序中所有中间值从设备拷贝到主机。

breakpoint 还支持 orderedtoken 参数(两者互斥):ordered=True 保证与其它 ordered 的打印/断点的相对顺序;token 则传入一个 JAX 数组(或数组的 PyTree),断点会在该值被计算后运行,token 原样返回并需接回后续计算,若返回值未被使用则整段计算会被剪枝、断点不触发。

优势:简单直观且(某种程度上)标准化;可同时检查调用栈上下的众多值。 局限:可能需要布设大量断点才能定位错误源头;物化大量中间值。


三、分片调试:inspect_array_shardingvisualize_sharding

除运行时数值调试外,jax.debug 还提供一组分片调试工具,用于在(以及不在)staged 函数内检视数组分片。

3.1 inspect_array_sharding

inspect_array_sharding(value, *, callback) 接受一个数组 PyTree,对其中每个数组以各自的 Sharding 调用回调,并能在 jax.jit 的计算中工作(jax/_src/debugging.py):

import jax
import jax.numpy as jnp
from jax.sharding import Mesh, PartitionSpec

x = jnp.arange(8, dtype=jnp.float32)

def f_(x):
  x = jnp.sin(x)
  jax.debug.inspect_array_sharding(x, callback=print)
  return jnp.square(x)

f = jax.jit(f_, in_shardings=PartitionSpec('dev'),
            out_shardings=PartitionSpec('dev'))

with jax.set_mesh(Mesh(jax.devices(), ('dev',))):
  f.lower(x).compile()
# NamedSharding(mesh={'dev': 8}, partition_spec=PartitionSpec(('dev',),))

回调触发的时机策略是"分片信息一可用就尽早调用"

  • 无任何变换时,数组与分片立即可得,回调立即执行;
  • jax.jit 内,分片在 lowering 阶段确定,可用 AOT API(jit(f).lower(...))触发;
  • jax.pjit 内,分片由 XLA 在编译期确定,可用 jax.jit(f).lower(...).compile() 触发;
  • 任何情况下运行函数都会触发(运行必然包含 lowering 与编译),但一旦函数编译完成并被缓存,回调不再发生

该 API 目前标记为实验性,行为未来可能变化。测试覆盖见 tests/debugging_primitives_test.py 中的 test_inspect_sharding_is_called_in_jittest_inspect_sharding_is_called_in_jit_shardedtest_inspect_sharding_3d_jit

3.2 visualize_array_shardingvisualize_sharding

visualize_array_sharding(arr, **kwargs)inspect_array_sharding 的可视化封装:它用 arr.shape 与检视到的分片调用 visualize_sharding,在终端中以富文本表格形式绘制分片布局(jax/_src/debugging.py)。

visualize_sharding(shape, sharding, ...) 的完整参数:

参数 含义
shape 数组形状,仅支持 1D 与 2D(其它维度抛 ValueError
sharding 要可视化的 Sharding 对象
use_color 是否启用颜色(默认 True,需终端支持颜色且已安装 matplotlib 提供 tab20b 色图;否则自动降级为无彩色)
scale 整体缩放比例,默认 1.0
min_width / max_width 单元格最小/最大宽度(默认 9 / 80)
color_map 自定义颜色映射函数

绘制原理:通过 sharding.devices_indices_map(shape) 计算每个设备对应的切片,再按分块位置组织成 rich.table 输出,每个单元格标注设备平台(如 GPUCPUTPU)与设备 id,行列尺寸按分块在整体形状中的占比缩放,直观反映 1D/2D 的切分方式。注意该函数要求安装 rich,否则抛出 ValueError。相关测试见 test_visualize_wide_arraytest_visualize_sharding_shard_map


四、底层实现:调试效果如何在 JAX 中流转

理解 jax.debug 的实现能帮你更好地预判其行为。核心实现在 jax/_src/debugging.py

  • 两个 Primitivedebug_callback_pdebug_print_pjax.debug.print 最终复用 debug_callback 的 lowering(debug_print_lowering_rule 内部直接调用 debug_callback_lowering)。
  • 两种效果类型DebugEffect(无序)与 OrderedDebugEffect(有序)。它们被注册为 lowerable、control-flow-allowed、remat-allowed、custom-derivatives-allowed 与 partial-eval-kept 效果(jax/_src/debugging.py),这意味着调试回调可以合法地存在于控制流、remat、自定义微分等变换内部——这正是"回调会被复制、丢弃或重排"的机制根源。
  • 变换规则
    • batching 规则将回调沿映射轴展开(unroll),即 vmap 下每个元素触发一次回调(jax/_src/debugging.py);
    • JVP 规则把回调留在 primal 路径、切线返回空,因此 jax.grad 下只在正向传播打印(jax/_src/debugging.py);
    • transpose 规则返回全 None,保证反向传播不会重复执行回调;
    • partial_eval 自定义规则尽可能把回调 staged 出来以提供更多信息(jax/_src/debugging.py)。
  • lowering 与分片策略:SPMD 上下文(shard_map/pjit)下,全自动分片使用 MAXIMAL sharding、回调只在逻辑整体值上执行一次;全手动分片则每个设备各执行一次;有序效果通过 token 串行化保证相对顺序。TPU 上的调试回调使用 channel id,因而注册为 cacheable=Falsejax/_src/debugging.py)。
  • 设备到主机的回传:回调最终通过 emit_python_callback 在主机侧执行 Python 函数(jax/_src/debugging.py),这也解释了"必须把设备值拷回主机"的性能开销来源。

这些行为在 tests/debugging_primitives_test.py 中有系统化验证,例如 test_can_stage_out_debug_printtest_debug_print_batchingtest_debug_print_jvp_ruletest_debug_print_grad_with_custom_vjp_ruletest_remat_of_debug_printtest_unordered_print_with_jittest_debug_print_two_call_form 等,可作为理解各变换语义的参考。


五、总结与选型建议

需求 推荐 API
jit/pmap/pjit 中打印 traced 数值 jax.debug.print(保持顺序用 ordered=True
需要自定义输出格式或副作用 jax.debug.callback
暂停执行、交互式检查调用栈 jax.debug.breakpoint(配合 jax.lax.cond 可做 NaN/Inf 断点)
检视中间值的分片信息 jax.debug.inspect_array_sharding
在终端可视化 1D/2D 分片布局 jax.debug.visualize_array_sharding / visualize_sharding(需 rich

使用 jax.debug 时务必牢记:调试语句会改变编译计算(可能引发融合差异与额外物化)、可能乱序/异步执行、并引入设备到主机的通信开销。它只服务于"揭示计算如何被求值"这一调试目的,不应出现在依赖语义保证的生产代码路径中。想系统了解 JAX 全部调试手段(含 checkify、调试 flag、XLA metadata 与慢编译排查),可继续阅读 docs/debugging/index.md

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

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
docsdocs
暂无描述
Markdown
900
5.83 K
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.76 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.94 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
603
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
527