Just-in-time compilationJAX JIT 编译原理与实战指南【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读jax.jit是 JAX 中最核心的变换之一它把一段 Python 函数在追踪tracing阶段规约成中间表示 jaxpr再交给 XLA 编译器做即时Just-In-Time编译最终生成针对 CPU/GPU/TPU 高度优化的可执行代码。本篇文章围绕 docs/jit-compilation.md 展开先讲透 JAX 变换与 jaxpr 的原理再手把手演示 SELU 算子从逐 op 执行到 JIT 加速的完整过程并深入剖析「为什么不能无脑 JIT 一切」、静态参数static_argnums/static_argnames的使用时机以及 JIT 缓存与重编译的行为边界。读完你将掌握 JAX JIT 的正确打开方式并能用jax.make_jaxpr观察函数内部到底发生了什么。JAX 变换是如何工作的JAX 允许我们变换 Python 函数其秘密在于JAX 会把每个函数规约成一系列 primitive原语操作每个 primitive 代表一个最基本的计算单元。而jax.make_jaxpr就是观察这一过程的窗口——它返回函数的 jaxprJAX 的中间表示让我们直观看到追踪结果。import jax import jax.numpy as jnp global_list [] def log2(x): global_list.append(x) ln_x jnp.log(x) ln_2 jnp.log(2.0) return ln_x / ln_2 print(jax.make_jaxpr(log2)(3.0))输出会是一段类似下面这样的 jaxpr{ lambda ; a:f32[]. let b:f32[] log a c:f32[] log 1.0 d:f32[] div b c in (d,) }关于 jaxpr 各字段lambda、类型、let 绑定、in 返回值的语义可以参考 docs/601/jaxpr.md 的详细说明。关键点副作用不会进入 jaxpr注意上面log2函数体里有global_list.append(x)但 jaxpr 里完全没有与它对应的内容。这并非 bug而是特性JAX 变换被设计为只理解**无副作用功能纯**的代码。如果对「纯函数」「副作用」这些术语不熟悉可以在仓库的 docs/notebooks/Common_Gotchas_in_JAX.md 的 Pure Functions 一节找到通俗解释。不纯函数在 JAX 变换下是危险的它们可能静默失败或者产生像Tracer 泄漏tracer leak那样令人困惑的下游错误而且 JAX 往往无法自动检测到副作用的存在。文档给出的官方替代方案是想要调试打印用jax.debug.print其实现位于 jax/_src/debug.py想表达通用副作用、接受性能损失用jax.experimental.io_callback实现在 jax/_src/callback.py 的io_callback想检查 tracer 泄漏、接受性能损失用jax.check_tracer_leaks。追踪tracing的微观机制追踪时JAX 会用tracer 对象包装每个参数。这些 tracer 会记录函数调用期间发生在普通 Python 层面对它们执行的所有 JAX 操作然后 JAX 用这些记录重建整个函数重建的产物就是 jaxpr。由于 tracer 不记录 Python 副作用副作用自然不会出现在 jaxpr 中——但注意副作用在追踪过程中仍然真实发生了。一个典型例子是 Python 的print()def log2_with_print(x): print(printed x:, x) ln_x jnp.log(x) ln_2 jnp.log(2.0) return ln_x / ln_2 print(jax.make_jaxpr(log2_with_print)(3.))你会发现打印出来的x是一个Traced对象——这正是 JAX 内部机制在起作用。「Python 代码至少会被执行一次」严格来说是实现细节不应依赖它但在调试时可以利用这一点打印中间值。jaxpr 只反映「参数给定」的那一次执行jaxpr 捕获的是函数在给定参数上的那次执行路径。如果函数里有 Python 条件分支jaxpr 只会包含实际走到的那个分支def log2_if_rank_2(x): if x.ndim 2: ln_x jnp.log(x) ln_2 jnp.log(2.0) return ln_x / ln_2 else: return x print(jax.make_jaxpr(log2_if_rank_2)(jax.numpy.array([1, 2, 3])))传入的是 1 维数组x.ndim 2为假所以 jaxpr 只会包含else分支的return x。这暗示了一个重要的推论jaxpr 的形状/类型是由追踪时的输入决定的一旦输入 shape 或 dtype 改变JAX 就不得不重新追踪和重新编译。从源码看jax.make_jaxpr的实现位于 jax/_src/api.py 的make_jaxpr函数约 L2136 起。它本质上是对jit(fun, static_argnumsstatic_argnums).trace(...)的一次封装先做一次纯追踪再把常量consts重新合并回 jaxpr 返回给调用者。其 docstring 也明确说明jaxpr 是「基于带 let 绑定的简单类型化一阶 lambda 演算」的追踪中间表示make_jaxpr返回的是抽象到ShapedArray级别的追踪结果。JIT 编译一个函数JAX 允许同一份代码跑在 CPU/GPU/TPU 上但默认逐 op 把操作发给加速器这限制了 XLA 编译器的优化空间。下面以深度学习常用的SELUScaled Exponential Linear Unit算子为例import jax import jax.numpy as jnp def selu(x, alpha1.67, lambda_1.05): return lambda_ * jnp.where(x 0, x, alpha * jnp.exp(x) - alpha) x jnp.arange(1000000) %timeit selu(x).block_until_ready()这段代码的问题在于每次只把单个 op 发给加速器XLA 无法看到全局做融合优化。我们的目标自然是把尽可能多的代码交给 XLA 编译器让它做整体优化。jax.jit就是为此设计的selu_jit jax.jit(selu) # 先预编译一次再计时 selu_jit(x).block_until_ready() %timeit selu_jit(x).block_until_ready()刚才发生了什么selu_jit jax.jit(selu)得到selu的编译版本一个被包装的函数。调用一次selu_jit(x)JAX 在这里做追踪——它必须有真实输入才能用 tracer 包装。得到的 jaxpr 再由 XLA 编译成针对 GPU/TPU 优化过的代码随后立即执行以满足这次调用。之后的每次调用都直接复用已编译代码完全跳过 Python 实现。如果不单独做 warm-up基准测试会把编译时间也计进去——虽然因为循环很多次整体仍会更快但那就不是公平对比了。计时对编译版本测速。注意这里用了block_until_ready()这是因为 JAX 采用异步派发async dispatch需要显式等待结果就绪再计时。关于异步派发机制可以参阅 docs/async_dispatch.rstblock_until_ready的定义见 jax/_src/api.py 与 jax/_src/array.py。值得说明的是jax.jit在 jax/_src/api.py 中的签名还支持更多参数完整的默认值如下节选in_shardings/out_shardings输入输出的分片sharding规格配合jax.sharding使用static_argnums/static_argnames标记静态编译期常量参数donate_argnums/donate_argnames标记可捐赠的缓冲区帮助 XLA 复用输入内存、降低峰值内存keep_unused默认False未被函数使用的参数会从编译产物中剔除、不传上设备device/backend显式指定运行设备或后端cpu/gpu/tpuinline嵌套 jitted 函数的内联策略默认False即jax.Inline.AUTOcompiler_options传递给 XLA 编译器的选项字典。其中donate_argnums是文档未展开但实践中很重要的能力捐赠后不能再复用这些缓冲区JAX 会在你尝试复用时报错。更多细节可以参考仓库中的 docs/buffer_donation.md。为什么不能无脑 JIT 一切看完上面的加速效果你可能会想干脆给所有函数都套上jax.jit得了。先来看两个 JIT 会失败的例子# 对 x 的「值」做条件判断 def f(x): if x 0: return x else: return 2 * x jax.jit(f)(10) # 报错 # 循环条件依赖 x 和 n 的值 def g(x, n): i 0 while i n: i 1 return x i jax.jit(g)(10, 20) # 报错根因用运行时值控制追踪期流程两个例子的共同问题是试图用运行时runtime值来控制追踪期trace-time的程序流程。在 JIT 内部被追踪的值如这里的x、n只能通过它们的静态属性——例如shape或dtype——来影响控制流而不能通过它们的具体数值。if x 0这样的判断发生在追踪期此时x是 tracer对它做布尔比较会直接抛出 tracer 错误。关于 Python 控制流与 JAX 的交互细节请参阅 docs/control-flow.md。解法一改写代码或使用 lax 控制流应对这个问题一种方式是改写代码、避免对值做条件判断另一种是使用 docs/201/control-flow.md 中介绍的特殊控制流原语比如jax.lax.cond其实现位于 jax/_src/lax/control_flow/conditionals.py。解法二只 JIT 编译函数的一部分有时候改写不现实那就考虑只 JIT 函数中计算最昂贵的部分。比如循环体是整个函数的热点就只 JIT 循环体但要小心下一节讲的缓存问题避免适得其反# 循环条件依赖 x 和 n但循环体是 JIT 的 jax.jit def loop_body(prev_i): return prev_i 1 def g_inner_jitted(x, n): i 0 while i n: i loop_body(i) return x i g_inner_jitted(10, 20)外层while i n仍是普通 Python 循环可以按值判断内层loop_body被 JIT 编译热点计算获得了加速。这是「部分 JIT」的典型模式。把参数标记为 static静态参数如果确实需要 JIT 一个「对输入值做条件判断」的函数可以告诉 JAX对某个输入使用抽象程度更低更具体的 tracer。方法是指定static_argnums按位置索引或static_argnames按参数名。代价是显著的静态参数的每个不同取值都会产生不同的 jaxpr 和编译产物JAX 不得不为每个新值重新编译。所以只有在该函数只会遇到有限的静态取值集合时这才是个好策略。f_jit_correct jax.jit(f, static_argnums0) print(f_jit_correct(10))g_jit_correct jax.jit(g, static_argnames[n]) print(g_jit_correct(10, 20))以装饰器形式使用时用装饰器工厂模式jax.jit(static_argnames[n]) def g_jit_decorated(x, n): i 0 while i n: i 1 return x i print(g_jit_decorated(10, 20))源码层面的约定与限制结合 jax/_src/api.py 中jit的 docstring静态参数还有一些容易踩坑的约定静态参数必须是可哈希的实现__hash__和__eq__且不可变因为它们的值会参与编译缓存键compilation cache key的计算。文档特别强调非数组类型或数组容器之外的参数必须标记为 static否则无法被正确追踪。如果只给了static_argnums而没给static_argnames或反之JAX 会用inspect.signature(fun)自动推断对应的参数名/位置如果两者都给了则只把显式列出的参数当作静态不再推断。从 JAX v0.8.1 起jit支持省略函数参数的装饰器工厂写法即jax.jit(static_argnames[n])而非partial(jax.jit, ...)上面的示例正是官方推荐的现代写法旧版本则常用functools.partial实现同样效果。JAX 对fun持有弱引用作为缓存键的一部分因此fun必须可被弱引用weakly-referenceable。JIT 与缓存第一次 JIT 调用有编译开销所以理解jax.jit何时、如何缓存编译结果是用好它的关键。缓存的基本规则假设f jax.jit(g)首次调用f时完成追踪 编译XLA 代码被缓存后续调用f直接复用缓存代码不再重复编译——这就是jax.jit摊平编译前期成本的方式。如果指定了static_argnums那么只有静态参数取值与缓存一致时才复用任何一个静态值变化都会触发重编译。如果静态参数取值范围很大程序花在编译上的时间可能比逐 op 执行还多——这是常见的性能陷阱。不要在循环里对临时函数调用 jit文档明确警告避免在循环或其他 Python 作用域内对临时函数调用jax.jit。原因在于缓存依赖函数的哈希当等价函数被反复重新定义哈希不同时缓存就失效了导致每次循环迭代都重新编译from functools import partial def unjitted_loop_body(prev_i): return prev_i 1 def g_inner_jitted_partial(x, n): i 0 while i n: # 别这么做每次 partial 返回的函数哈希都不同 i jax.jit(partial(unjitted_loop_body))(i) return x i def g_inner_jitted_lambda(x, n): i 0 while i n: # 别这么做lambda 每次也返回哈希不同的函数 i jax.jit(lambda x: unjitted_loop_body(x))(i) return x i def g_inner_jitted_normal(x, n): i 0 while i n: # 这样没问题JAX 能找到缓存的编译函数 i jax.jit(unjitted_loop_body)(i) return x i print(jit called in a loop with partials:) %timeit g_inner_jitted_partial(10, 20).block_until_ready() print(jit called in a loop with lambdas:) %timeit g_inner_jitted_lambda(10, 20).block_until_ready() print(jit called in a loop with caching:) %timeit g_inner_jitted_normal(10, 20).block_until_ready()结论很直观partial和lambda每次都会产生新的函数对象、新的哈希缓存形同虚设而直接传入同一个模块级函数时JAX 能稳定命中缓存。缓存机制的底层佐证从源码角度看JAX 在派发层大量使用带缓存的装饰器来复用编译产物例如 jax/_src/dispatch.py 中的xla_primitive_callableL98 附近使用util.cache()缓存 primitive 的 callable并有util.test_event(xla_primitive_callable_cache_miss)这样的测试探针标记缓存未命中事件同一文件还有多处util.weakref_lru_cache/util.cache(max_size2048, ...)用于缓存各类派发中间结果。这正是「同一函数对象反复jit能命中缓存、哈希不同的临时函数会反复编译」这一文档结论在实现层面的体现。小结JIT 的正确打开方式尽量 JIT 大片代码把尽可能多的计算交给 XLA让它做算子融合与设备级优化用block_until_ready()配合%timeit得到不含异步派发误差的公平基准。保持函数纯追踪期只记录 primitive 操作副作用既不进 jaxpr 也会带来 tracer 泄漏等隐患调试打印用jax.debug.print通用副作用用jax.experimental.io_callback。控制流按「静态属性」而非「值」走需要按值分支/循环时改写代码、改用jax.lax.cond等原语或只 JIT 热点内层部分。静态参数要克制static_argnums/static_argnames只适用于取值集合有限、可哈希的场景否则会陷入重编译泥潭。把jit用在稳定、可缓存的函数对象上避免在循环内用partial/lambda现造函数再 JIT。如果需要更系统地学习可以继续阅读仓库中 docs/201/jit.md 关于jit进阶语义、docs/601/jaxpr.md 关于 jaxpr 语言以及 docs/async_dispatch.rst 关于异步派发的说明。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
