
用 JAX 的思维方式编程从 NumPy 到 Tracing、JIT 与 XLA 的完整实战指南【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jaxJAX 提供了一套简单而强大的 API 用于编写加速数值计算代码但要高效使用 JAX往往需要理解它底层的运行机制。本篇指南以 JAX 官方教程 thinking_in_jax.md 为骨架结合本仓库源码jax/_src/api.py、jax/_src/core.py、jax/_src/numpy/lax_numpy.py等逐层剖析 JAX 与 NumPy 的异同、jax.numpy/jax.lax/ XLA 三层 API 结构、Tracing 与静态变量的内部原理以及jax.jit的适用边界与最佳实践。读完后你将能够判断一段代码该用numpy还是jax.numpy、何时需要 JIT、为什么某些操作无法被 JIT 编译并能熟练运用make_jaxpr、static_argnums等工具排查和编写高性能 JAX 程序。JAX 与 NumPy接口相似语义不同本节核心概念JAX 提供了一套灵感来源于 NumPy 的接口便于上手。借助 Python 的鸭子类型duck-typingJAX 数组在多数场景下可以作为 NumPy 数组的直接替身使用。与 NumPy 数组不同JAX 数组永远不可变immutable。熟悉而强大的jax.numpy接口NumPy 提供了广为人知、功能强大的数值数据 API。为方便用户JAX 提供了jax.numpy它高度镜像 NumPy 的 API是进入 JAX 世界最轻松的入口。几乎所有能用numpy完成的工作都可以用jax.numpy完成import matplotlib.pyplot as plt import numpy as np x_np np.linspace(0, 10, 1000) y_np 2 * np.sin(x_np) * np.cos(x_np) plt.plot(x_np, y_np);import jax.numpy as jnp x_jnp jnp.linspace(0, 10, 1000) y_jnp 2 * jnp.sin(x_jnp) * jnp.cos(x_jnp) plt.plot(x_jnp, y_jnp);上面两段代码除了将np换成jnp之外完全相同绘制结果也一致。可以看到JAX 数组在很多场景下例如绘图可以直接顶替 NumPy 数组使用。不过两者底层实现是不同类型的 Python 对象type(x_np) # numpy.ndarray type(x_jnp) # jax.ArrayPython 的鸭子类型让 JAX 数组与 NumPy 数组在大量场景下可以互换。从源码看Tracer 类 与 jax.Array 都实现了__array__、__dlpack__等协议这正是互操作性的基础。JAX 数组不可变用.at[]进行索引更新JAX 数组与 NumPy 数组之间有一个重要差异JAX 数组是不可变的一旦创建其内容就不能被原地修改。下面先看 NumPy 中修改数组的例子# NumPy: mutable arrays x np.arange(10) x[0] 10 print(x) # [10 1 2 3 4 5 6 7 8 9]而 JAX 中的等价操作会直接报错# JAX: immutable arrays x jnp.arange(10) x[0] 10 # TypeError: ... jax.Array object does not support item assignment要更新单个元素JAX 提供了索引更新语法indexed update返回一个更新后的新副本原数组保持不变y x.at[0].set(10) print(x) # [0 1 2 3 4 5 6 7 8 9] print(y) # [10 1 2 3 4 5 6 7 8 9]在仓库源码中这套 API 由 array_methods.py 中的_IndexUpdateRef类 实现除了set之外还提供了add、multiply、divide、min、max等变体见 set/add/multiply 方法并且支持indices_are_sorted、unique_indices等性能提示参数。不可变性是 JAX 函数式语义的基石也是其能够安全地执行 JIT、自动微分等变换的前提。NumPy、lax 与 XLAJAX 的 API 分层本节核心概念jax.numpy是高层封装提供熟悉的接口。jax.lax是更底层、更严格、通常也更强大的 API。所有 JAX 操作最终都归结为 XLAAccelerated Linear Algebra 编译器中的原语操作。jax.lax更严格也更强大的底层 API查看jax.numpy的源码会发现其中所有操作最终都是用jax.lax中定义的函数来表达的。可以把jax.lax看作一个更严格、但通常也更强大的多维数组 API。最典型的差异体现在**类型提升type promotion**上jax.numpy会隐式提升参数类型以允许混合数据类型的运算而jax.lax不会import jax.numpy as jnp jnp.add(1, 1.0) # jax.numpy API 隐式提升混合类型正常返回 array(2., dtypefloat32)from jax import lax lax.add(1, 1.0) # jax.lax API 要求显式类型提升报 TypeError直接使用jax.lax时遇到这类情况需要显式完成类型提升lax.add(jnp.float32(1), 1.0) # array(2., dtypefloat32)这一“严格性”并非缺陷它让类型行为可预测这对编译器和自动微分都至关重要。仓库源码可以佐证在 ufuncs.py 中的add实现 中jnp.add本质上就是lax.add对于 bool 类型则退化为lax.bitwise_or可见jax.numpy是jax.lax上的一层薄封装。从jnp.convolve到lax.conv_general_dilated在严格之外jax.lax还提供了一些比 NumPy 更通用的高效操作。例如考虑一维卷积NumPy 风格的写法是x jnp.array([1, 2, 1]) y jnp.ones(10) jnp.convolve(x, y) # Array([1., 3., 4., 4., 4., 4., 4., 4., 4., 4., 3., 2.], dtypefloat32)在底层这个 NumPy 操作会被翻译为lax.conv_general_dilated这个通用得多的卷积原语from jax import lax result lax.conv_general_dilated( x.reshape(1, 1, 3).astype(float), # 注意这里需要显式类型提升 y.reshape(1, 1, 10), window_strides(1,), padding[(len(y) - 1, len(y) - 1)]) # 等价于 NumPy 的 paddingfull result[0, 0] # Array([1., 3., 4., 4., 4., 4., 4., 4., 4., 4., 3., 2.], dtypefloat32)这是一个面向深度神经网络常用卷积场景设计的批量卷积操作。它需要更多样板代码但远比 NumPy 提供的卷积更灵活、更具扩展性更深入的讲解见仓库中的 convolutions.md。值得注意的一个实现细节是本仓库中 lax_numpy.py 的convolve实现 本身就用partial(jit, static_argnames(mode, precision, preferred_element_type))做了 JIT 封装内部调用lax.conv_general_dilated完成计算支持full、same、valid三种模式以及precision、preferred_element_type等参数。这生动地展示了jax.numpy函数既是“lax 的封装”又是“JIT 的受益者”。一切操作最终都归于 XLA从根本上看所有jax.lax操作都是 XLA 中相应操作的 Python 包装例如上面的卷积实现对应 XLA 的ConvWithGeneralPadding语义。JAX 的每个操作最终都会表达为这些基础 XLA 操作这正是 JIT 编译得以实现的根基——因为整段计算可以被整体翻译成一份 XLA 计算图HLO交给编译器统一优化。To JIT or not to JIT何时编译、为何编译本节核心概念默认情况下JAX 一次只执行一个操作op-by-op 模式按顺序运行。使用即时编译JIT装饰器后一长串操作可以被一起优化并一次性执行。并非所有 JAX 代码都能被 JIT 编译因为 JIT 要求数组形状在编译期已知且为静态值。默认执行模式与jax.jit变换由于所有 JAX 操作都表达为 XLA 操作JAX 可以借助 XLA 编译器非常高效地执行代码块。例如下面这个用jax.numpy操作对二维矩阵的行做标准化的函数import jax.numpy as jnp def norm(X): X X - X.mean(0) return X / X.std(0)可以用jax.jit变换创建它的即时编译版本from jax import jit norm_compiled jit(norm)编译后的函数与原始函数在标准浮点精度内返回相同结果np.random.seed(1701) X jnp.array(np.random.rand(10000, 10)) np.allclose(norm(X), norm_compiled(X), atol1E-6) # True由于编译带来了操作融合fusing、避免临时数组分配以及一系列其他优化JIT 编译后的执行时间可能比逐算子模式快几个数量级注意下面用block_until_ready()来保证计时准确性这与 JAX 的异步分发机制有关%timeit norm(X).block_until_ready() %timeit norm_compiled(X).block_until_ready()为什么需要block_until_ready()在 array.py 的实现 中可以看到该方法会等待底层所有设备缓冲区的计算真正完成确保计时反映的是实际计算耗时而不是异步派发返回的时间。关于异步分发机制的详细讨论见仓库中的 async_dispatch.rst。JIT 的边界静态形状要求当然jax.jit也有其限制它要求所有数组具有静态形状static shapes。这意味着部分 JAX 操作与 JIT 编译不兼容。例如下面这个操作可以在逐算子模式下正常执行def get_negatives(x): return x[x 0] x jnp.array(np.random.randn(10)) get_negatives(x)但若尝试在 jit 模式下执行就会报错jit(get_negatives)(x) # TracerArrayConversionError / ConcretizationTypeError原因在于该函数生成的结果数组形状在编译期不可知输出的大小取决于输入数组的值因此与 JIT 不兼容。理解“哪些值在编译期可知”是掌握 JIT 的关键这正是下一节的主题。从源码看jax.jit的完整签名本仓库 api.py 中的jit定义 给出了完整的参数签名除了文档重点介绍的static_argnums/static_argnames外还包括in_shardings/out_shardings指定输入输出的分片Sharding约束用于分布式场景donate_argnums/donate_argnames指定哪些参数缓冲区可以被计算“捐赠”覆盖复用帮助 XLA 减少内存分配keep_unused是否保留函数未使用的参数默认False未用参数不会被传输到设备device/backend指定运行的设备或后端cpu、gpu、tpuinline是否将该函数内联进外层 jaxpr默认False。这些参数使得jax.jit不仅是性能工具也是控制内存与分布式行为的入口。JIT 机制Tracing 与静态变量本节核心概念JIT 及其他 JAX 变换通过**跟踪tracing**一个函数来确定它对特定形状与类型的输入会产生什么效果。你不想被跟踪的变量可以标记为静态static。用print观察 tracer要高效使用jax.jit有必要理解它的工作机制。在 JIT 编译的函数中加入几个print()语句再调用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 语句确实执行了但打印出来的并不是我们传入的数据而是tracer 对象——它们是真实数据的占位符。tracer 正是jax.jit提取函数操作序列的机制。基础的 tracer 只编码数组的形状shape和数据类型dtype对具体值不敏感。这段被记录下来的计算序列随后可以在 XLA 中高效地应用于任何同形状、同 dtype 的新输入而无需重新执行 Python 代码。在仓库源码中tracer 的基类是 core.py 中的Tracer类可以看到其shape、dtype、ndim、size等属性都来自_aval_property即由抽象值AbstractValue简称 aval派生——这印证了“tracer 只关心形状与类型”这一结论。同时tracer 上调用tolist()、__array__、__dlpack__等方法会抛出ConcretizationTypeError这正是“tracer 不携带具体值”这一约束在代码层面的体现。当我们再次以形状匹配的输入调用编译后的函数时不需要重新编译也不会再打印任何内容——因为结果由编译后的 XLA 计算得到而不是 Pythonx2 np.random.randn(3, 4) y2 np.random.randn(4) f(x2, y2)jaxpr被记录下来的计算表达式被提取的操作序列以 JAX 表达式jaxprJAX expression 的缩写编码。可以用jax.make_jaxpr变换查看from jax import make_jaxpr def f(x, y): return jnp.dot(x 1, y 1) make_jaxpr(f)(x, y)make_jaxpr在 api.py 中定义它接受static_argnums、axis_env、abstracted_axes等参数返回一个基于示例输入产生 jaxpr 的函数也可通过return_shapeTrue同时返回输出形状。jaxpr 是 JAX 各种变换如jax.jit、jax.grad、jax.vmap之间传递的中间表示是整个体系的中枢。控制流不能依赖 traced 值由上述机制可以推出一个重要结论因为 JIT 编译不包含数组的具体值信息函数内的控制流语句不能依赖 traced 值。例如下面这段代码会报错jit def f(x, neg): return -x if neg else x f(1, True) # ConcretizationTypeError: Attempted to convert value to a Python int如果有些变量你不想被跟踪可以将它们标记为静态参数from functools import partial partial(jit, static_argnums(1,)) def f(x, neg): return -x if neg else x f(1, True) # Array(-1, dtypeint32, weak_typeTrue)注意以不同的静态参数调用 JIT 函数会触发重新编译因此函数依然能按预期工作f(1, False) # Array(1, dtypeint32, weak_typeTrue)从 api.py 中jit的文档 可知静态参数会作为编译缓存键的一部分参与缓存寻址因此它们必须是可哈希实现__hash__与__eq__且不可变的对象非数组类或容器类的参数则必须标记为静态。理解哪些值与操作会被静态处理、哪些会被跟踪是高效使用jax.jit的关键能力。静态操作与 Traced 操作本节核心概念正如值可以分为静态与 traced 两类操作也可以分为静态与 traced 两类。静态操作在编译期由 Python 求值traced 操作被编译进 XLA在运行时求值。想要操作是静态的用numpy想要操作被跟踪用jax.numpy。一个常见的“踩坑”案例静态值与 traced 值的区分让我们必须思考如何保持一个静态值始终静态。考虑这个函数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)它会报错错误信息大意是期望一个整数类型的具体值序列却发现了 tracer。往函数里加一些 print 语句来定位原因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是 traced 的但x.shape是一个静态值。然而一旦用jnp.array和jnp.prod处理这个静态值它就变成了traced 值于是不能再用于reshape()这类要求静态输入的函数中回忆一下数组形状必须是静态的。numpy与jax.numpy的正确分工这里有一个非常实用的模式静态操作在编译期完成用numpytraced 操作编译进运行时用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) # Array([[1., 1., 1., 1., 1., 1.]], dtypefloat32)正是出于这个原因JAX 程序中的标准惯例是同时import numpy as np和import jax.numpy as jnp从而对每个操作都能精确控制它是在编译期静态执行numpy只运行一次还是在运行时被跟踪执行jax.numpy随计算图被优化。小结构建 JAX 思维方式的四个要点回顾整篇指南从 NumPy 平滑过渡到高效 JAX 编程需要内化以下四层认识接口层jax.numpy提供了与 NumPy 几乎一致、且可互操作的接口但数组不可变原地修改要换成.at[]索引更新语法实现见 array_methods.py层次结构jax.numpy→jax.lax→ XLA 三层递进jax.lax更严格类型提升需显式但更通用所有操作最终都表达为 XLA 原语执行模式默认逐算子执行jax.jit可以将整段计算融合编译从而大幅提速但要求形状静态可知测量性能时要结合异步分发机制使用block_until_ready()Tracing 模型JIT 通过 tracer 记录计算图jaxpr控制流不能依赖 traced 值必要时用static_argnums标记静态参数并用numpy/jax.numpy分别控制静态与 traced 操作。这套思维方式不仅适用于jax.jit也适用于jax.grad、jax.vmap等其他所有 JAX 变换——它们共享同一套 tracing 基础设施。掌握了它你就掌握了写出正确、高效、可扩展 JAX 代码的通用方法论。本教程对应的可执行 Notebook 版本位于仓库中的 thinking_in_jax.ipynb配套的入门教程还包括 quickstart.md、autodidax.md从头实现 JAX 的核心机制以及 How JAX primitives work可作为进一步深入的学习路线。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考