JAX Tracing 机制详解:理解 jit 如何把 Python 函数编译为可执行计算
JAX Tracing 机制详解理解 jit 如何把 Python 函数编译为可执行计算【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxdocs/tracing.md是 JAX 官方教程中关于tracing追踪机制的权威讲解。本文以该教程为骨架结合本仓库 jax/_src/core.py 与 jax/_src/api.py 的源码实现系统性地剖析jax.jit等变换如何通过 tracer 提取计算图、生成 jaxpr以及静态值 vs 追踪值抽象 tracer vs 具体 tracer这两组核心概念对控制流、形状运算和编译性能的实际影响。读完本文你将能够诊断常见的 TracerArrayConversionError 与 ConcretizationTypeError 报错并正确使用static_argnums、numpy与jax.numpy的组合来编写可高效编译的 JAX 代码。什么是 Tracingjax.jit的工作方式jax.jit以及其他 JAX 变换transform的工作方式是**追踪trace**一个函数确定它对某个特定形状与类型的输入会产生什么效果。为了直观理解 tracing最直接的办法是在一个被 JIT 编译的函数里放入几条print()语句然后调用该函数from jax import jit import jax.numpy as jnp import numpy as np jit def f(x, y): print(Running f():) print(f x {x}) print(f y {y}) result jnp.dot(x 1, y 1) print(f result {result}) return result x np.random.randn(3, 4) y np.random.randn(4) f(x, y)观察输出可以发现print语句确实执行了但打印出来的并不是我们传入的真实数据而是用于**替身stand-in**的 tracer 对象——形如TracedShapedArray(float32[])这样的内容。这正是jax.jit提取函数所描述的操作序列sequence of operations的方式基本的 tracer 是替身编码了数组的 shape 与 dtype但对具体的数值不敏感agnostic to the values。这段被记录下来的计算序列随后可以在 XLA 中高效地应用到具有相同 shape 与 dtype 的新输入上而无需重新执行 Python 代码。从源码看tracer 就是 jax/_src/core.py 中定义的Tracer类每个 tracer 实例都携带一个avalabstract value抽象值属性class TracerTraceType: Trace: __slots__ [__weakref__, _trace, _line_info, aval] ... aval: AbstractValue dtype _aval_property(dtype) ndim _aval_property(ndim) size _aval_property(size) shape _aval_property(shape)shape、dtype、ndim、size这些属性全部从aval派生——这正是tracer 只关心形状与类型、不关心数值这一说法的源码级体现。第二次调用无需重新编译当我们再次以匹配的输入调用编译后的函数时无需重新编译也不会打印任何内容因为结果是在编译好的 XLA 计算中算出来的而不是在 Python 中x2 np.random.randn(3, 4) y2 np.random.randn(4) f(x2, y2) # 无输出走的是已编译的 XLA 计算这一行为意味着 tracing 只在**首次调用或缓存未命中时**发生一次之后的热路径hot path完全由编译产物接管。jaxprTracing 的产物tracing 提取出的操作序列被编码为一种 JAX 表达式即jaxprJAX expression 的缩写参见 key-concepts.md 中关于 jaxprs 的章节以及 601/jaxpr.md 对 jaxpr 语言的完整介绍。可以使用jax.make_jaxpr变换查看 jaxprfrom jax import make_jaxpr def f(x, y): return jnp.dot(x 1, y 1) make_jaxpr(f)(x, y)make_jaxpr返回一个包装后的函数将示例参数应用于该函数后即可得到函数在这些参数上的 jaxpr 表示。从 jax/_src/api.py 的 docstring 可知Ajaxpris JAXs intermediate representation for program traces. Thejaxprlanguage is based on the simply-typed first-order lambda calculus with let-bindings. ... Thejaxprreturned is a trace offunabstracted toShapedArraylevel.即jaxpr 是基于简单类型一阶 lambda 演算含 let 绑定的程序追踪中间表示make_jaxpr返回的 jaxpr 是函数在ShapedArray抽象层级上的追踪结果。此外make_jaxpr还支持static_argnums与jax.jit的参数含义一致指定哪些位置参数作为静态值处理return_shapeTrue返回值变为(jaxpr, pytree)二元组其中第二个元素是与函数输出同结构的、每个叶子带shape与dtype属性的表示可用于获取输出类型信息。ShapedArray本身就是 jax/_src/core.py 中定义的抽象值类其核心字段正是shape、dtype、weak_type、sharding、memory_space、layout——全部是元数据层面而非数值层面的信息进一步印证了抽象值的含义。控制流不能依赖被追踪的值tracing 有一个直接后果需要特别注意因为 JIT 编译是在不了解数组具体内容的情况下进行的函数中的控制流语句不能依赖被追踪的值详见 control-flow.md。例如下面的代码会报错jit def f(x, neg): return -x if neg else x f(1, True)原因在于if neg需要对neg求布尔值而 JIT 传入的neg是一个 tracertracer 本身不携带具体数值因此无法在 Python 层决定走哪个分支。在源码层面tracer 上需要具体值才能完成的操作如bool()、tolist()、tobytes()、__dlpack__()以及访问device、committed等属性都会抛出ConcretizationTypeError见 jax/_src/core.py这正是tracer 不能被当作具体值使用这一规则的强制保证。用 static_argnums 保持静态如果有些变量不希望被追踪可以把它们标记为静态static使其成为编译期常量from functools import partial jit(static_argnums(1,)) def f(x, neg): return -x if neg else x f(1, True)注意以不同的静态参数调用 JIT 编译的函数会触发重新编译所以函数行为依然正确f(1, False) # 触发重新编译静态参数值从 True 变为 False从 jax/_src/api.py 中jax.jit的签名与 docstring 可以看到static_argnums的完整语义指定哪些位置参数作为trace- 和 compile-time constant追踪期与编译期常量处理静态参数必须可哈希同时实现__hash__与__eq__且不可变否则可以是任意 Python 对象以不同的值调用这些常量会触发重新编译非数组类型或非数组容器的参数必须标记为静态与static_argnames配合使用时若只提供一个JAX 会通过inspect.signature(fun)自动推导对应的另一个若两者都提供则不再使用inspect.signature只把明确列出的参数视为静态。static_argnums还支持传入整数序列如(0, 2)从而一次标记多个参数。静态操作 vs 被追踪的操作与值可以是静态或追踪的相对应操作operation同样可以是静态或追踪的静态操作在编译期于 Python 中求值追踪操作被编译并在运行期于 XLA 中求值。理解这种静态 vs 追踪的区分关键在于思考如何保持一个静态值不被污染为追踪值。考虑下面的函数import jax.numpy as jnp from jax import jit jit def f(x): return x.reshape(jnp.array(x.shape).prod()) x jnp.ones((2, 3)) f(x) # 报错reshape 需要静态的 1D 整数序列却得到了 tracer该函数会报错错误信息指出本应传入一维的具体整数序列1D sequence of concrete values of integer type却发现了 tracer。通过加打印语句可以弄清来龙去脉jit def f(x): print(fx {x}) print(fx.shape {x.shape}) print(fjnp.array(x.shape).prod() {jnp.array(x.shape).prod()}) # return x.reshape(jnp.array(x.shape).prod()) # 注释掉以避开报错 f(x)观察输出可以发现虽然x是被追踪的但x.shape是静态值它是从 tracer 的aval.shape属性直接读取的元数据属于 jax/_src/core.py 中_aval_property定义的派生属性整个求值过程不产生任何新操作。然而一旦对x.shape使用jnp.array和jnp.prod这个静态值就被卷入追踪世界变成了追踪值从而无法再用于像reshape()这样要求静态输入的函数回忆数组形状必须是静态的。实用模式numpy 管静态jax.numpy 管追踪一个非常有用的模式是对于应该静态的即在编译期完成操作使用numpy对于应该被追踪的即编译后在运行期执行的操作使用jax.numpy。对上面的函数修正写法是from jax import jit import jax.numpy as jnp import numpy as np jit def f(x): return x.reshape((np.prod(x.shape),)) f(x) # 正常工作np.prod(x.shape)在 Python 层立即计算出一个普通整数形状元组的乘积reshape于是拿到了静态的形状参数。正因为如此JAX 程序中的标准惯例是同时import numpy as np与import jax.numpy as jnp以便精细控制每个操作是在静态方式numpy编译期执行一次还是追踪方式jax.numpy运行期优化执行下完成。判断哪些值和操作会是静态的、哪些会被追踪是高效使用jax.jit的关键能力。以下是两条实用的判别准则情形结果从 tracer 读取shape、dtype、ndim、size等元数据属性得到静态Python 值对静态值调用jnp.*中的函数如jnp.array、jnp.prod被污染为追踪值对静态值调用np.*中的函数如np.prod保持静态用静态值做reshape、slicing、控制流条件等正常工作用追踪值做上述需要具体值的操作抛出ConcretizationTypeError/TracerArrayConversionError不同种类的 JAX 值抽象 tracer 与具体 tracertracer 值携带一个抽象值abstract value例如携带形状与 dtype 信息的ShapedArray——这样的 tracer 称为抽象 tracerabstract tracer。但并非所有 tracer 都是抽象的有些 tracer例如自动微分变换为参数引入的 tracer携带的抽象值是ConcreteArray其中实际包含常规的数组数据可用于解析条件分支——这样的 tracer 称为具体 tracerconcrete tracer。由具体 tracer可能与常规值组合计算出的 tracer 值其结果仍是具体 tracer。而**具体值concrete value**则指常规值或具体 tracer。一般来说只要计算涉及至少一个 tracer 值其结果就是 tracer 值。但有极少数例外——当计算可以完全借助 tracer 携带的抽象值完成时结果可以是常规 Python 值获取携带ShapedArray抽象值的 tracer 的shape显式把具体 tracer 值转换为常规类型例如int(x)或x.astype(float)bool(x)当具体性允许时会产生 Python bool——这个情形在控制流中尤其重要因为它极其常见。各变换引入何种 tracer下表总结了各变换对位置参数引入 tracer 的具体规则对应原教程的权威表格变换引入的 tracer 类型例外jax.jit所有位置参数引入抽象 tracerstatic_argnums指定的参数保持为常规值jax.pmap所有位置参数引入抽象 tracerstatic_broadcasted_argnums指定的参数除外jax.vmap、jax.make_jaxpr、xla_computation所有位置参数引入抽象 tracer无jax.jvp、jax.grad所有位置参数引入具体 tracer当处于外层变换中、实际参数本身已是抽象 tracer 时自动微分引入的 tracer 也变为抽象 tracer高阶控制流原语lax.cond、lax.while_loop、lax.fori_loop、lax.scan处理函数体functional时引入抽象 tracer无论当前是否在进行 JAX 变换实战基于数据的条件控制流以上内容在你编写只能作用于常规 Python 值的代码例如基于数据的条件控制流时非常关键。看这个例子def divide(x, y): return x / y if y 1. else 0.如果要对它应用jax.jit必须指定static_argnums1确保y保持为常规值。原因是布尔表达式y 1.需要具体值常规值或 tracer 都行但 JIT 默认引入的是抽象 tracer不满足要求。同理如果显式写bool(y 1.)、int(y)或float(y)也会遇到同样的问题。有意思的是jax.grad(divide)(3., 2.)可以直接工作因为jax.grad引入的是具体 tracer条件分支可以利用y的具体数值来解析。这正是具体 tracer 携带真实数据、可用于解析条件分支这一特性的直接体现也解释了为什么grad能处理jit处理不了的数据相关控制流。调试与进阶把 tracing 变成可视化的工具结合以上原理当你遇到与 tracing 相关的报错时可以按下面的思路快速定位看报错类型TracerArrayConversionError试图把 tracer 转成 numpy 数组、TracerBoolConversionError试图对 tracer 求布尔值、ConcretizationTypeErrortracer 上访问了需要具体值的属性/方法——这三类错误都源于把追踪值当具体值用在函数内加print打印 tracer 会显示其ShapedArray(shape, dtype)摘要可以快速确认哪些值是静态的、哪些是被追踪的例如x.shape是静态的而jnp.array(x.shape)是追踪的用make_jaxpr检查计算图确认追踪到的是否符合预期、是否有意外引入的额外操作保持静态纯度凡是参与reshape、切片的起始/终止索引、if/else条件等必须是具体值位置的量一律用numpy在编译期算好凡是参与数值运算的量一律用jax.numpy。在 jax/_src/core.py 中可以看到tracer 的aval是唯一的抽象值来源shape等元数据通过_aval_property从aval派生而make_jaxprjax/_src/api.py则以ShapedArray抽象层级对函数进行追踪并返回 jaxpr。理解了抽象值 shape dtype 其他元数据这一事实也就从根本上理解了为什么 JIT 编译出的代码与具体数据内容无关以及为什么它能在新输入上被反复复用而无需重新编译。小结Tracing 是jax.jit等变换的工作基础用编码了 shape/dtype 的 tracer 替身执行函数提取操作序列提取出的操作序列编码为jaxpr可用jax.make_jaxpr查看该 jaxpr 是在ShapedArray抽象层级上的函数追踪结果由于编译期不知道数组内容控制流不能依赖被追踪的值需要用static_argnums或static_argnames把相关参数标记为静态静态操作在 Python/编译期执行追踪操作在 XLA/运行期执行用numpy保持静态、用jax.numpy保持追踪是编写可编译 JAX 代码的核心惯例tracer 分抽象 tracer只带 shape/dtype与具体 tracer携带真实数据两类jit/vmap/pmap引入抽象 tracergrad/jvp引入具体 tracerlax.cond等控制流原语引入抽象 tracer——这一差异决定了代码中的条件分支在哪些变换下可用、哪些不可用。延伸阅读控制流的完整处理方式见 control-flow.mdjaxpr 语言的深入说明见 601/jaxpr.md 与 601/jax-primitives.mdJIT 编译的更多细节见 jit-compilation.md。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →