)
JAX jaxpr 语言深度解析读懂 trace 产生的内部中间表示IR【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxJAX 之所以能在很小的代码体积内同时支持自动微分、向量化与 JIT 编译核心秘密在于它用 Python 解释器本身完成了降维把任意 Python NumPy 程序蒸馏成一种简单、静态类型化、一阶的中间表示语言——jaxpr。本文以 docs/601/jaxpr.md 为骨架结合仓库中jax/_src/core.py、jax/_src/api.py的源码实现系统讲解 jaxpr 的语法结构、jax.core.Jaxpr的数据模型、jax.make_jaxpr的使用方法以及cond、while、scan、jit等携带子 jaxpr 的高阶原语。读完本文你将能够亲手用make_jaxpr观察任意函数的 trace 结果并准确解读其中每一行方程的含义。jaxpr 是什么JAX 的中间表示IRJaxpr 是 JAX 程序内部的中间表示Intermediate Representation, IR。它具备四个关键性质显式类型explicitly typed、函数式functional、一阶first-order以及代数正规形式ANF, Algebraic Normal Form。从概念上讲可以认为 JAX 的各种变换如jax.jit、jax.grad、jax.vmap遵循这样的流程首先对要被变换的 Python 函数做trace 特化trace-specializing将其转化为一种小型、行为良好的中间形式然后用变换特有的解释规则去解释这个中间形式。JAX 强大之处在于它从一个熟悉且灵活的编程接口Python NumPy出发利用真实的 Python 解释器完成大部分提炼工作把计算的核心提炼成一种几乎没有高阶特性、静态类型化的表达式语言——jaxpr 语言。需要特别注意的是并非所有 JAX 变换都会逐字构造出 jaxpr。有些变换例如自动微分grad、批处理vmap会在 trace 过程中增量地施加变换并不一定物化一个完整 jaxpr。但如果你想理解 JAX 内部如何工作或者想利用 JAX trace 的产物例如调试、序列化、导出理解 jaxpr 就是必经之路。仓库佐证在 docs/601/index.rst 中jaxpr 被定位为 JAX 内部机制系列的第一课——tracing produces 的中间表示、它的文法以及如何读懂它面向对 JAX 内部好奇的贡献者与底层扩展者。jaxpr 的语法术语term语法jaxpr 术语语法如下jaxpr :: { lambda binder , ... . let eqn ... in ( atom , ... ) } binder :: var:array_type var :: a | b | c | ... atom :: var | literal literal :: int32 | int64 | float32 | float64 eqn :: binder , ... primitive [ params ] atom , ...可以看出jaxpr 是一个显式绑定变量的 let 表达式lambda声明输入参数let定义一系列方程equationin列出输出表达式。方程只依赖输入变量和此前方程定义出的中间变量因此天然满足 ANF 形式。每个变量都带有数组类型标注binder 中的var:array_type体现显式类型特性。并非所有 Python 程序都能被这种形式处理但事实证明绝大多数科学计算和机器学习程序都可以。打印printed语法jax.make_jaxpr打印出的 jaxpr 使用如下文法jaxpr :: { lambda Var* ; Var. let Eqn* in [Expr] }其中分号两侧的两组变量是 jaxpr 的参数第一组Var*分号之前是为被提升hoist出来的常量引入的变量称为constvars在ClosedJaxpr中consts字段保存了对应的值。第二组Var分号之后称为invars对应被 trace 的 Python 函数的输入。Eqn*是方程列表每个方程用一个原语primitive作用在一些原子表达式上定义一个或多个中间变量。每个方程只能使用输入变量和之前方程定义的中间变量。Expr是 jaxpr 的输出原子表达式列表字面量或变量。方程Equation打印如下Eqn :: let Var Primitive [ Param* ] Expr其中Var是一个或多个中间变量作为一次原语调用的输出有些原语会返回多个值。Expr是一个或多个原子表达式每个要么是变量、要么是字面量常量。特殊变量unitvar或字面量unit打印为*代表该值在后续计算中不再需要、已被省略——它只是占位符。Param*是零个或多个传给原语的命名参数打印在方括号[]中每个参数形如Name Value。绝大多数 jaxpr 原语是一阶的只接受一个或多个Expr作为参数Primitive : add | sub | sin | mul | ...最常见的 jaxpr 原语在jax.lax模块中有系统文档对应仓库中的 jax.lax.rst。代码数据模型jax.core.Jaxpr与ClosedJaxpr一个 jaxpr 实例表示一个带一个或多个类型化参数输入变量、一个或多个类型化结果的函数。结果只依赖于输入变量不存在从外层作用域捕获的自由变量。输入和输出都带有类型在 JAX 中类型用抽象值abstract values表示。代码中有两个相关的 jaxpr 表示jax.core.Jaxpr和jax.core.ClosedJaxpr。ClosedJaxpr表示一个部分应用的Jaxpr也就是jax.make_jaxpr返回的对象包含jaxpr一个jax.core.Jaxpr表示函数实际的下述计算内容consts常量列表。ClosedJaxpr最有趣的部分正是Jaxpr所承载的实际执行内容。从当前仓库源码看jax/_src/core.pyclass Jaxpr的实现细节如下通过__slots__保存_all_invars、_outvars、_eqns、_effects、_debug_info、_is_high、_consts等字段all_invars属性返回全部输入变量constvars invars 拼接在一起见构造函数中self._all_invars [*constvars, *invars]constvars属性是带有附着值的前缀输入变量即self._all_invars[: len(self._consts)]——这与文档中constvars 与 invars 的区别只是簿记惯例的说法完全一致invars属性返回去掉常量前缀之后的输入变量outvars返回输出变量列表eqns返回方程列表effects返回副作用集合构造函数还支持旧式ClosedJaxpr(jaxpr, consts)兼容调用并保留了jaxpr属性作为 legacy 访问器源码注释明确写着 Legacy accessor from the days of ClosedJaxpr, which wrapped a Jaxpr——也就是说在新版本中ClosedJaxpr已经被合并进Jaxpr类jaxpr属性直接返回自身。每条方程在源码中对应class JaxprEqnjax/_src/core.py其字段与打印语法一一对应invars: list[Atom]—— 方程输入原子表达式outvars: list[Var]—— 方程输出变量primitive: Primitive—— 所用的原语params: dict[str, Any]—— 命名参数打印在方括号中的Name Valueeffects: Effects—— 该方程可能产生的副作用source_info与ctx—— 源码位置等调试信息。公共 API 提示底层实现位于内部模块jax/_src/core.py其公共入口是jax.extend.core可参考 jax.extend.core.rst 与 docs/601/jax-primitives.md 中from jax.extend import core的用法。用jax.make_jaxpr观察第一个 jaxprjax.make_jaxpr返回一个给定示例参数即可得到 jaxpr的函数。从源码看jax/_src/api.py其完整签名为make_jaxpr(fun, static_argnums(), axis_envNone, return_shapeFalse)static_argnums将指定位置的参数视为静态不可 trace参数用法与jax.jit一致return_shape为True时返回(jaxpr, shape)元组axis_env轴环境供涉及命名轴的变换使用。来看文档中的第一个例子from jax import make_jaxpr import jax.numpy as jnp def func1(first, second): temp first jnp.sin(second) * 3. return jnp.sum(temp) print(make_jaxpr(func1)(jnp.zeros(8), jnp.ones(8)))输出示意{ lambda ; a:f32[8] b:f32[8]. let c:f32[8] sin b d:f32[8] mul c 3.0 e:f32[8] add a d f:f32[] reduce_sum[axes(0,) input_shape(8,)] e in (f,) }解读如下这里没有 constvarsa、b是输入变量分别对应first、second两个函数参数标量字面量3.0直接内联在方程里标量常量不需要提升为 constvar见后文常量变量一节reduce_sum原语除了操作数e之外还带有命名参数axes(0,)和input_shape(8,)以[Name Value]的形式打印在方括号内。这个例子直观展示了打印语法中invars、Eqn*、Expr、Param*的对应关系。Python 控制流与函数调用会被内联重要的一点是即使执行一个调用 JAX 的程序会构建 jaxprPython 级别的控制流和 Python 级别的函数调用仍然会照常执行。因此Python 程序里含有函数和控制流并不意味着生成的 jaxpr 必须包含控制流或高阶特性。例如trace 下面的func3时JAX 会把对inner的调用以及if second.shape[0] 4这个条件判断完全内联产生与之前func1完全相同的 jaxprdef func2(inner, first, second): temp first inner(second) * 3. return jnp.sum(temp) def inner(second): if second.shape[0] 4: return jnp.sin(second) else: assert False def func3(first, second): return func2(inner, first, second) print(make_jaxpr(func3)(jnp.zeros(8), jnp.ones(8)))由于jnp.zeros(8)和jnp.ones(8)的 shape 是静态已知的shape 为 8满足 4if在 trace 期间即被解析、else分支的assert False根本不会进入。最终 jaxpr 与func1完全一致——这就是Python 控制流发生在 trace 时而非运行时的典型体现。这也引出了一个核心使用原则如果希望控制流在运行时动态执行例如循环次数取决于运行时的数组值就必须显式使用jax.lax.cond、jax.lax.while_loop等构造它们会在 jaxpr 中留下高阶原语见下文。处理 pytrees元组被展平在 jaxpr 中没有元组类型原语接受多个输入、产生多个输出。当被处理函数带有结构化输入或输出时JAX 会把它们展平flatten在 jaxpr 中呈现为输入/输出列表。有关展平的完整机制可参考仓库中的 pytrees 教程。例如下面的代码产生与前面func1完全相同的 jaxpr两个输入变量对应输入元组的两个元素def func4(arg): # The arg is a pair. temp arg[0] jnp.sin(arg[1]) * 3. return jnp.sum(temp) print(make_jaxpr(func4)((jnp.zeros(8), jnp.ones(8))))这说明无论 Python 层的输入是两个位置参数还是一个二元组经过 pytree 展平后jaxpr 层面看到的都是同样的一组输入变量。常量变量constvarsjaxpr 中的某些值是与参数无关的常量。标量常量直接内联在方程中如前面例子里的3.0非标量的数组常量则被提升到 jaxpr 顶层成为常量变量constvars。constvars 与其他 jaxpr 参数invars的唯一区别只是簿记惯例——在ClosedJaxpr中consts字段保存着与这些 constvars 一一对应的值。源码层面同样印证了这一点Jaxpr.constvars的实现就是带有附着值的前缀输入变量jax/_src/core.py构造函数中_all_invars [*constvars, *invars]表明两组变量共用同一个存储仅靠前缀长度区分num_consts。当你在子 jaxpr如cond的 branches 或while的 body中捕获外层数组常量时它就会以 constvar 的形式出现在该子 jaxpr 的lambda分号之前。高阶 JAX 原语除了普通的一阶原语jaxpr 还包含若干高阶higher-orderJAX 原语。它们更复杂因为其参数中嵌入了子 jaxprsub-jaxpr。cond原语条件分支JAX 会 trace 普通的 Python 条件语句。若要将条件表达式捕获为运行时动态执行必须使用jax.lax.switch和jax.lax.cond构造器签名如下lax.switch(index: int, branches: Sequence[A - B], operand: A) - B lax.cond(pred: bool, true_body: A - B, false_body: A - B, operand: A) - B两者在内部都会绑定一个名为cond的原语。jaxpr 中的cond原语反映了更一般的lax.switch签名它接受一个整数表示要执行的分支索引会被钳制到合法的索引范围内。例如from jax import lax def one_of_three(index, arg): return lax.switch(index, [lambda x: x 1., lambda x: x - 2., lambda x: x 3.], arg) print(make_jaxpr(one_of_three)(1, 5.))输出示意{ lambda ; a:i32[] b:f32[]. let c:f32[] cond[ branches( ...子 jaxpr 1..., ...子 jaxpr 2..., ...子 jaxpr 3... ) linear(False,) ] a b in (c,) }cond原语的参数branches与各分支函数对应的 jaxpr。本例中每个分支函数都接受一个输入变量对应xlinear一个布尔元组由自动微分机制内部使用编码哪些输入参数在条件中被线性使用。上述cond实例接受两个操作数第一个打印为d是分支索引第二个b是传给branches中被选中 jaxpr 的操作数即arg。再看使用jax.lax.cond的例子from jax import lax def func7(arg): return lax.cond(arg 0., lambda xtrue: xtrue 3., lambda xfalse: xfalse - 3., arg) print(make_jaxpr(func7)(5.))此时布尔谓词被转换为整数索引0 或 1branches中的 jaxpr 依次对应 false 分支与 true 分支函数注意顺序false 在前。同样每个函数接受一个输入变量分别对应xfalse与xtrue。再看一个更复杂的情况分支函数的输入是元组且 false 分支函数内部含有常量jnp.ones(1)——它会被提升为 constvardef func8(arg1, arg2): # Where arg2 is a pair. return lax.cond(arg1 0., lambda xtrue: xtrue[0], lambda xfalse: jnp.array([1]) xfalse[1], arg2) print(make_jaxpr(func8)(5., (jnp.zeros(1), 2.)))这里你能在输出中清楚地看到branches里的 false 分支 jaxpr 以lambda ; a:f32[1]之外多出一组 constvar 开头jnp.array([1])被提升这正是常量变量一节所述行为的真实案例。while原语循环与条件分支类似Python 循环在 trace 期间会被内联。若要在运行时动态执行循环必须使用jax.lax.while_loop原语或jax.lax.fori_loop生成 while_loop 原语的辅助函数lax.while_loop(cond_fun: (C - bool), body_fun: (C - C), init: C) - C lax.fori_loop(start: int, end: int, body: (int - C - C), init: C) - C其中C表示循环carry值的类型。示例import numpy as np def func10(arg, n): ones jnp.ones(arg.shape) # A constant. return lax.fori_loop(0, n, lambda i, carry: carry ones * 3. arg, arg ones) print(make_jaxpr(func10)(np.ones(16), 5))while原语共接受 5 个参数示意输出中为c a 0 b d0 个cond_jaxpr的常量因为cond_nconsts为 02 个body_jaxpr的常量即c和a——ones与arg被捕获为 body 中的 constvar3 个 carry 初始值的参数打印为0 b d之类对应arg ones展平后的若干输入。fori_loop在此处实际上等价于一个带索引计数器的while_loop循环次数n是运行时整数因此不能静态展开只能以while原语动态执行。scan原语静态长度的数组循环JAX 支持一种对数组元素进行循环的特化形式其迭代次数在编译期静态已知。正因迭代次数固定这种循环可以方便地做反向微分reverse-differentiable。这类循环用jax.lax.scan构造lax.scan(body_fun: (C - A - (C, B)), init_carry: C, in_arr: Array[A]) - (C, Array[B])其中C是 scan carry 的类型A是输入数组的元素类型B是输出数组的元素类型。示例函数func11def func11(arr, extra): ones jnp.ones(arr.shape) # A constant def body(carry, aelems): # carry: running dot-product of the two arrays # aelems: a pair with corresponding elements from the two arrays ae1, ae2 aelems return (carry ae1 * ae2 extra, carry) return lax.scan(body, 0., (arr, ones)) print(make_jaxpr(func11)(np.ones(16), 5.))scan原语的linear参数描述每个输入变量是否被保证在 body 中被线性使用一旦 scan 经过线性化linearization更多参数会变为线性——这与cond原语中的linear参数目的一致都服务于自动微分的内部记账。scan原语共接受 4 个参数示意输出中为b 0.0 a c1 个 body 的自由变量extra被捕获进 body 子 jaxpr1 个 carry 的初始值0.02 个 scan 操作的数组arr与ones对应a、c。(p)jit原语callcall 原语源自 JIT 编译它封装一个子 jaxpr并附带指定计算运行后端backend与设备device的参数。示例from jax import jit def func12(arg): jit def inner(x): return x arg * jnp.ones(1) # Include a constant in the inner function. return arg inner(arg - 2.) print(make_jaxpr(func12)(1.))输出示意{ lambda ; a:f32[]. let b:f32[1] pjit[ nameinner jaxpr{ lambda ; c:f32[] d:f32[1]. let e:f32[1] mul c d f:f32[1] add c e in (f,) } ... ] a g:f32[] add a b in (g,) }这里inner被jit包裹tracefunc12时生成 call 原语对应jit/pjit其参数中包含一个子 jaxprarg标量与jnp.ones(1)常量分别作为 invar 与 constvar 进入子 jaxpr。外层 jaxpr 通过该子 jaxpr 组织出对inner的调用。从源码角度看call 类原语的这一模式也普遍存在于仓库各变换解释器中子 jaxpr 作为一种嵌套的、可被单独编译/变换的单元正是高阶原语实现组合变换的基础。若想了解原语需要为 JAX 各变换impl / abstract_eval / lowering / JVP / transpose / batching提供哪些规则请继续阅读 docs/601/jax-primitives.md。总结与延伸阅读jaxpr 语言是理解 JAX 内部机制的第一块基石它是显式类型、函数式、一阶、ANF的中间表示由lambda/let/in三部分构成用jax.make_jaxpr可以观察任意 Python 函数的 trace 结果输入 pytree 会被展平、Python 控制流与函数调用会被内联、标量常量内联而数组常量提升为 constvar动态控制流cond、while与静态长度循环scan、JIT 调用call/pjit以高阶原语的形式出现在 jaxpr 中其参数中嵌入子 jaxpr数据模型上Jaxpr及其前身ClosedJaxpr现已在 jax/_src/core.py 中合并通过constvars/invars/outvars/eqns/consts等字段精确对应上述打印语法JaxprEqn则承载单条方程的原语、参数与副作用信息。如果希望更深入地掌握 JAX 内部可以按 docs/601/index.rst 的路线继续阅读 jax-primitives 了解原语如何支撑各类变换再跟随 autodidax 教程 从零用纯 Python 一步步实现 tracing、jaxpr、自动微分与 jit最终完整复现 JAX 的核心设计。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考