JAX jaxpr 语言深度解析:读懂 trace 产生的内部中间表示(IR)
2026/9/10 20:26:03 网站建设 项目流程

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.pyjax/_src/api.py的源码实现,系统讲解 jaxpr 的语法结构、jax.core.Jaxpr的数据模型、jax.make_jaxpr的使用方法,以及condwhilescanjit等携带子 jaxpr 的高阶原语。读完本文,你将能够亲手用make_jaxpr观察任意函数的 trace 结果,并准确解读其中每一行方程的含义。

jaxpr 是什么:JAX 的中间表示(IR)

Jaxpr 是 JAX 程序内部的中间表示(Intermediate Representation, IR)。它具备四个关键性质:显式类型(explicitly typed)函数式(functional)一阶(first-order)以及代数正规形式(ANF, Algebraic Normal Form)

从概念上讲,可以认为 JAX 的各种变换(如jax.jitjax.gradjax.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.JaxprClosedJaxpr

一个 jaxpr 实例表示一个带一个或多个类型化参数(输入变量)、一个或多个类型化结果的函数。结果只依赖于输入变量,不存在从外层作用域捕获的自由变量。输入和输出都带有类型,在 JAX 中类型用"抽象值(abstract values)"表示。

代码中有两个相关的 jaxpr 表示:jax.core.Jaxprjax.core.ClosedJaxprClosedJaxpr表示一个"部分应用"的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_infoctx—— 源码位置等调试信息。

公共 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,ab是输入变量,分别对应firstsecond两个函数参数;
  • 标量字面量3.0直接内联在方程里(标量常量不需要提升为 constvar,见后文"常量变量"一节);
  • reduce_sum原语除了操作数e之外,还带有命名参数axes=(0,)input_shape=(8,),以[Name = Value]的形式打印在方括号内。

这个例子直观展示了打印语法中invarsEqn*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.condjax.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.switchjax.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 在前)。同样,每个函数接受一个输入变量,分别对应xfalsextrue

再看一个更复杂的情况:分支函数的输入是元组,且 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的常量(即ca——onesarg被捕获为 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 操作的数组(arrones,对应ac)。

(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,) }

这里innerjit包裹,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;
  • 动态控制流(condwhile)与静态长度循环(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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询