JAX jaxpr 语言深度解析:读懂 trace 产生的内部中间表示(IR)
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
JAX 之所以能在很小的代码体积内同时支持自动微分、向量化与 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 的中间表示(IR)
Jaxpr 是 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定义一系列方程(equation),in列出输出表达式。方程只依赖输入变量和此前方程定义出的中间变量,因此天然满足 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.py),class 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 JaxprEqn(jax/_src/core.py),其字段与打印语法一一对应:
invars: list[Atom]—— 方程输入原子表达式;outvars: list[Var]—— 方程输出变量;primitive: Primitive—— 所用的原语;params: dict[str, Any]—— 命名参数(打印在方括号中的Name = Value);effects: 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观察第一个 jaxpr
jax.make_jaxpr返回一个"给定示例参数即可得到 jaxpr"的函数。从源码看(jax/_src/api.py),其完整签名为:
make_jaxpr(fun, static_argnums=(), axis_env=None, return_shape=False)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,) }解读如下:
- 这里没有 constvars,
a、b是输入变量,分别对应first、second两个函数参数; - 标量字面量
3.0直接内联在方程里(标量常量不需要提升为 constvar,见后文"常量变量"一节); reduce_sum原语除了操作数e之外,还带有命名参数axes=(0,)和input_shape=(8,),以[Name = Value]的形式打印在方括号内。
这个例子直观展示了打印语法中invars、Eqn*、Expr+、Param*的对应关系。
Python 控制流与函数调用会被内联
重要的一点是:即使执行一个调用 JAX 的程序会构建 jaxpr,Python 级别的控制流和 Python 级别的函数调用仍然会照常执行。因此,Python 程序里含有函数和控制流,并不意味着生成的 jaxpr 必须包含控制流或高阶特性。
例如,trace 下面的func3时,JAX 会把对inner的调用以及if second.shape[0] > 4这个条件判断完全内联,产生与之前func1完全相同的 jaxpr:
def 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,满足> 4),if在 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 层面看到的都是同样的一组输入变量。
常量变量(constvars)
jaxpr 中的某些值是与参数无关的常量。标量常量直接内联在方程中(如前面例子里的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-order)JAX 原语。它们更复杂,因为其参数中嵌入了子 jaxpr(sub-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。本例中每个分支函数都接受一个输入变量(对应x);linear:一个布尔元组,由自动微分机制内部使用,编码哪些输入参数在条件中被线性使用。
上述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 或 1),branches中的 jaxpr 依次对应 false 分支与 true 分支函数(注意顺序:false 在前)。同样,每个函数接受一个输入变量,分别对应xfalse与xtrue。
再看一个更复杂的情况:分支函数的输入是元组,且 false 分支函数内部含有常量jnp.ones(1)——它会被提升为 constvar:
def 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 d):
- 0 个
cond_jaxpr的常量(因为cond_nconsts为 0); - 2 个
body_jaxpr的常量(即c和a——ones与arg被捕获为 body 中的 constvar); - 3 个 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是输出数组的元素类型。
示例函数func11:
def 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 c):
- 1 个 body 的自由变量(
extra,被捕获进 body 子 jaxpr); - 1 个 carry 的初始值(
0.0); - 2 个 scan 操作的数组(
arr与ones,对应a、c)。
(p)jit原语(call)
call 原语源自 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[ name=inner 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),其参数中包含一个子 jaxpr;arg(标量)与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 Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考