JAX 运行时调试实战指南:深入掌握 `jax.debug` 模块的打印、断点与分片检查
jax.debug 是 JAX 提供的一组运行时调试工具,专门用于在 jax.jit、jax.pmap、jax.pjit 等延迟执行(staged)的变换中打印数值、暂停执行并检查调用栈中的中间值,以及检视数组的分片(sharding)信息。读完本文,你将掌握 jax.debug.print、jax.debug.callback、jax.debug.breakpoint 三大运行时值调试 API 与 inspect_array_sharding、visualize_array_sharding、visualize_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):
callback、print、breakpoint,以及内部使用的效果标记DebugEffect、OrderedDebugEffect与log。官方完整指南见 docs/debugging/print_breakpoint.md。 - 分片调试(Sharding debugging):
inspect_array_sharding、visualize_array_sharding、visualize_sharding,用于在(以及不在)staged 函数内部检视和可视化数组分片。
全部实现位于 jax/_src/debugging.py(打印、回调与分片可视化)与 jax/_src/debugger/(断点调试器)中,jax/debug.py 只是薄薄的转发层。
一、jax.debug.print:在编译函数内打印数值
1.1 为什么普通 print 在 jax.jit 中失效
JAX 的变换分为两类:jax.grad、jax.vmap 会保留 Python 层的即时执行语义,此时内置 print 能打印数值;而 jax.jit、jax.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}")。从源码看,_DebugPrintFormatChecker(jax/_src/debugging.py)会主动检查未使用的参数并抛出错误提示 "You may be passing an f-string (i.e, f\"{x}\")...",帮助你尽早发现误用。
1.3 何时使用 jax.debug.print
- 需要打印动态(被 trace 的)数组值时,在
jit、vmap等变换内使用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 内置 print 在 grad 下的行为一致,但 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 同样支持两种调用形式(两段式推荐),参数 ordered、partitioned 含义与 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=True 在 jax.pmap 等并行变换下会报错,因为并行执行无法保证顺序。从源码看,ordered 对应 OrderedDebugEffect 效果类型,在 MLIR lowering 时会通过 token 机制串行化(jax/_src/debugging.py)。
计算扰动(Computation perturbation)
加入 jax.debug.print 或 jax.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.) # ==> 执行到断点时暂停!
与 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 还支持 ordered 与 token 参数(两者互斥):ordered=True 保证与其它 ordered 的打印/断点的相对顺序;token 则传入一个 JAX 数组(或数组的 PyTree),断点会在该值被计算后运行,token 原样返回并需接回后续计算,若返回值未被使用则整段计算会被剪枝、断点不触发。
优势:简单直观且(某种程度上)标准化;可同时检查调用栈上下的众多值。 局限:可能需要布设大量断点才能定位错误源头;物化大量中间值。
三、分片调试:inspect_array_sharding 与 visualize_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_jit、test_inspect_sharding_is_called_in_jit_sharded 与 test_inspect_sharding_3d_jit。
3.2 visualize_array_sharding 与 visualize_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 输出,每个单元格标注设备平台(如 GPU、CPU、TPU)与设备 id,行列尺寸按分块在整体形状中的占比缩放,直观反映 1D/2D 的切分方式。注意该函数要求安装 rich,否则抛出 ValueError。相关测试见 test_visualize_wide_array 与 test_visualize_sharding_shard_map。
四、底层实现:调试效果如何在 JAX 中流转
理解 jax.debug 的实现能帮你更好地预判其行为。核心实现在 jax/_src/debugging.py:
- 两个 Primitive:
debug_callback_p与debug_print_p,jax.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)。
- batching 规则将回调沿映射轴展开(unroll),即
- lowering 与分片策略:SPMD 上下文(
shard_map/pjit)下,全自动分片使用 MAXIMAL sharding、回调只在逻辑整体值上执行一次;全手动分片则每个设备各执行一次;有序效果通过 token 串行化保证相对顺序。TPU 上的调试回调使用 channel id,因而注册为cacheable=False(jax/_src/debugging.py)。 - 设备到主机的回传:回调最终通过
emit_python_callback在主机侧执行 Python 函数(jax/_src/debugging.py),这也解释了"必须把设备值拷回主机"的性能开销来源。
这些行为在 tests/debugging_primitives_test.py 中有系统化验证,例如 test_can_stage_out_debug_print、test_debug_print_batching、test_debug_print_jvp_rule、test_debug_print_grad_with_custom_vjp_rule、test_remat_of_debug_print、test_unordered_print_with_jit 与 test_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。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0634
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown00
jforgamejforgame是一个一站式游戏服务器开发框架。包含游戏服务器开发所需要的各种组件,比如网关,socket服务端与客户端,自定义高效消息编解码,游戏热更新,游戏通用工具等等。包含游戏服,跨服,匹配服,后台管理系统等实现,同时提供大量业务案例以供学习。亦可用于其他socket应用,例如及时聊天等。Java01
fizz-gateway-nodeAn Aggregation API Gateway in Java . FizzGate 是一个基于 Java开发的微服务聚合网关,是拥有自主知识产权的应用网关国产化替代方案,能够实现热服务编排聚合、自动授权选择、线上服务脚本编码、在线测试、高性能路由、API审核管理、回调管理等目的,拥有强大的自定义插件系统可以自行扩展,并且提供友好的图形化配置界面,能够快速帮助企业进行API服务治理、减少中间层胶水代码以及降低编码投入、提高 API 服务的稳定性和安全性。Java00
certd开源SSL证书管理工具;全自动证书申请、更新、续期;通配符证书,泛域名证书申请;证书自动化部署到阿里云、腾讯云、主机、群晖、宝塔;https证书,pfx证书,der证书,TLS证书,nginx证书自动续签自动部署JavaScript00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
