JAXjax.extend.core深入解析:Jaxpr 中间表示、Primitive 原语与底层核心机制
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
jax.extend.core是 JAX 为扩展者开放的"内部机制库视图",集中导出了 Jaxpr 中间表示(IR)、Primitive 原语、Effect 效果系统、类型抽象与追踪(tracing)相关的全部核心符号。本文以 docs/jax.extend.core.rst 所列 API 为骨架,逐一对齐 jax/extend/core/init.py、jax/_src/core.py、jax/_src/effects.py 等源码实现,帮助读者掌握"如何定义新的 JAX 原语"以及"如何读取、构造、校验 jaxpr"这两条最核心的扩展路径。
一、jax.extend与jax.extend.core的定位
1.1 什么是jax.extend
jax.extend模块(定义见 jax/extend/init.py)由 JEP #15856(对应仓库文档 docs/jep/15856-jex.md)提出,目的是把 JAX 的"内部组件"整理成一个二级 API 库视图,供 JAX 生态中的下游库(如 Oryx、jax-triton 以及各类自定义变换/编译器前端)使用。其关键定位是:
- 无兼容性保证:
jax.extend不遵循公开 API 的兼容性策略,不承诺弃用窗口,也不承诺跨版本向后兼容,破坏性变更通过 CHANGELOG.md 公布; - 区别于
jax.experimental:experimental是新特性的试验场,最终可能进入其他模块或被移除;而extend是从jax._src等内部包"搬迁"出来的稳定化符号集合; - 典型受众:需要写自定义变换、自研自动微分系统、编译器前端,或深度依赖 JAX IR 的开发者。
1.2jax.extend.core提供什么
JEP 对jax.extend.core的设想是:让调用方至少能够定义新的 JAX primitive,并能够处理jax.make_jaxpr产生的 Jaxpr IR,具体包括:
- 访问核心系统原语(如
add_p、mul_p等); - 访问 IR 类型(
Jaxpr、JaxprEqn、Var、Literal等); - 用于检查和格式化打印 jaxpr 的函数(
check_jaxpr等); - 显式构造 jaxpr 的工具(
new_jaxpr_eqn、gensym等)。
当前仓库中,jax.extend.core的符号分两部分导出:
- 主要符号从
jax._src.core(核心 IR 与追踪机制)、jax._src.abstract_arrays(array_types)等模块直接再导出,见 jax/extend/core/init.py; - 预注册好的系统原语(
primitives子模块)汇集在 jax/extend/core/primitives.py,其中包含了add_p、dot_general_p、while_p、psum_p、qr_p等数百个_p结尾的原语句柄。
注意:
jax.extend.core中的符号基本都标注为"不稳定",其中部分函数名带_DO_NOT_USE后缀,表明它们是出于兼容性保留的内部实现细节,新代码不应依赖。
二、Jaxpr:JAX 的核心中间表示
Jaxpr(JAX Expression)是 JAX 变换的核心中间表示。jax.make_jaxpr把 Python 函数"跟踪"(trace)后得到的就是一纸 Jaxpr。jax.extend.core提供了一整套用于表示和操作 Jaxpr 的类型。
2.1Jaxpr与ClosedJaxpr
Jaxpr类定义在 jax/_src/core.py,是整棵 IR 树的根。其核心字段(全部以只读 property 暴露)为:
| 属性 | 含义 |
|---|---|
all_invars | 全部输入变量(含常量变量),类型为list[Var] |
constvars | 常量输入(all_invars中前len(consts)个带值的输入) |
consts | 常量参数值列表(literals是它的历史别名) |
invars | 真正的输入变量(扣除常量后) |
outvars | 输出变量/常量列表,类型为list[Atom] |
eqns | 方程序列,类型为list[JaxprEqn] |
effects | 该 jaxpr 的整体效果集合(Effects) |
debug_info | 调试信息(DebugInfo) |
is_high | 是否包含高层原语(高/低两层 IR 机制的一部分) |
in_avals/out_avals | 输入/输出的抽象值列表 |
值得注意的实现细节:ClosedJaxpr与Jaxpr在源码中已经合并为同一个类。在 jax/_src/core.py 中可以看到:
# ClosedJaxpr and Jaxpr have been merged into a single class: a Jaxpr carries # a possibly-empty list of constant argument values, `consts`. ClosedJaxpr = JaxprClosedJaxpr保留为别名,供仍然使用ClosedJaxpr(jaxpr, consts)构造方式或isinstance判断的旧调用方使用(jaxprproperty 与map_jaxpr、replace(jaxpr=..., consts=...)等也是为兼容旧接口保留的)。因此今天"闭式 jaxpr"(不依赖外部自由变量、可直接执行的 jaxpr)与普通Jaxpr是同一类型,区别仅在于是否附带consts常量值。
此外Jaxpr提供pretty_print()(支持source_info、print_shapes、print_effects等开关)用于可读化输出,并在 IPython 中通过_repr_pretty_支持彩色美化打印。
2.2JaxprEqn:一条方程
JaxprEqn定义在 jax/_src/core.py,表示 jaxpr 中的一条指令,等价于"把某个 primitive 应用到若干输入、产出若干输出":
invars: list[Atom]:输入(Var或Literal);outvars: list[Var]:输出变量(多输出原语如eigh_p时是多个);primitive: Primitive:对应的原语;params: dict[str, Any]:编译期静态参数(如dot_general的维度约定);effects: Effects:本条方程的效果;source_info:源代码位置与 name stack;ctx: JaxprEqnContext:方程创建时的求值上下文快照(compute_type、抽象 mesh、布局模式、xla_metadata 等,见 jax/_src/core.py)。
源码注释特别说明JaxprEqn刻意用带__slots__的普通类而不是 NamedTuple,因为构造方程是热路径,需要更快的速度。它提供replace()方法,用于在改写 IR 时生成部分字段更新的副本。
2.3Var、Literal、DropVar与Atom
Var(core.py):jaxpr 中的变量节点,__slots__只有一个字段aval(其抽象值)。repr形如Var(id=...):f32[3,4];gensym(core.py):一个工厂函数gensym = lambda: Var,用来在构造 jaxpr 时生成"全新的变量类型"(配合Literal语义,通常用于中间变量占位);Literal(core.py):常量节点,持有val(值)与aval(抽象值)两个字段。只有literalable_types/literalable_scalar_types中登记的类型(如 NumPy 标量、TypedNdArray、Python 标量,见 jax/_src/abstract_arrays.py)才会被嵌入 jaxpr 成为Literal,其余常量会被提升为constvars;DropVar(core.py):Var的特例,表示"该赋值结果永远不会被读取"(被丢弃的输出),美化打印为_;Atom:类型别名Atom = Var | Literal(core.py),用于标注"方程输入/输出既可以来自变量也可以来自常量"的位置。
关于常量还有一组配套工具:is_literalable(x, for_ad=False)判断值能否作为Literal内联(for_ad=True时保留常量以便自动微分);is_hoistable(v)判断非标量大常量是否需要提升为函数参数(阈值由配置embedded_constants_max_bytes控制);jaxpr_const_args(jaxpr)收集 jaxpr 中需要提升的非标量常量。这些细节与 docs/internals/constants.md 中描述的常量机制一一对应。
2.4 jaxpr 遍历与子 jaxpr
jaxprs_in_params(params)(core.py):以生成器方式遍历方程params字典中的全部Jaxpr值(params中单个值或元组中任一元素是Jaxpr都会产出);subjaxprs(jaxpr)(core.py):产出jaxpr.eqns中所有子 jaxpr(不递归下钻,其 docstring 明确说明这一点)。while_p、scan_p、cond_p等控制流原语都以"把子 jaxpr 放在 params 里"的方式工作,这两函数是遍历嵌套 IR 的基础工具。
三、Primitive:一切运算的最小单元
3.1 Primitive 类结构
Primitive定义在 jax/_src/core.py,是 JAX 中最底层的运算描述符。每个 primitive 有:
name:字符串名称;- 若干类级标志位:
multiple_results(多输出)、call_primitive(final-style 调用原语)、ref_primitive(引用原语)、skip_canonicalization、ref_allocating; - 可被注册的实现与规则槽位:
impl(直接执行)、abstract_eval(抽象求值)、bind(绑定并分派到当前 trace)、bind_with_trace、get_bind_params、to_lojax等。
3.2bind:分派核心
当用户代码调用prim.bind(*args, **params)(core.py)时发生的关键流程:
- 对每个参数做规范化(
dtypes.canonicalize_value)并计算其抽象值aval; - 若参数是失效的 Tracer(逃逸出变换作用域),抛出
escaped_tracer_error; - 取出当前 trace(
prev_trace = trace_ctx.trace),临时将全局 trace 置空,再调用bind_with_trace; bind_with_trace(core.py)判断:若self.is_high(*avals, **params)且当前 trace 要求低层(requires_low),则走to_lojax下变换;否则调用trace.process_primitive(self, args, params)——这就是"在 jit 下生成方程、在 eager 下直接执行"的分派点。
is_high的默认实现(core.py)检查 params 中是否含有is_high=True的子 jaxpr,这是 JAX 高/低两层 IR(high-level / low-level jaxpr)机制的入口之一。
3.3 规则注册
Primitive提供一组def_*便捷方法用于注册各阶段规则:
| 方法 | 作用 |
|---|---|
def_impl(impl) | 注册直接求值实现(impl默认抛NotImplementedError) |
def_abstract_eval(abstract_eval) | 注册无副作用抽象求值规则(自动补no_effects作为第二返回值) |
def_effectful_abstract_eval(effectful_abstract_eval) | 注册带效果的抽象求值规则 |
def_effectful_abstract_eval2(abstract_eval) | 注册效果由GenericEffect(prim)概括的抽象求值规则 |
def_bind_with_trace(...) | 覆盖默认的分派行为 |
def_transpose/def_jvp/def_batching等 | 注册各变换规则(分别在 ad / ad / batching 解释器中约定,不在 core 类上) |
_effect_free_abstract_eval(core.py)把无副作用规则包装成返回(out, no_effects)二元组的形式——这正是现代 JAX 中abstract_eval统一返回"输出抽象值 + 效果集合"约定的一部分。
3.4primitives子模块:系统原语全集
jax/extend/core/primitives.py 把散布在jax._src各处的系统原语统一再导出,是"想要复用 JAX 内置运算语义"时的总索引,主要分组包括:
- 基础逐元素运算:
add_p、mul_p、sub_p、div_p、pow_p、exp_p、log_p、sin_p、tanh_p等; - 形状/切片:
reshape_p、transpose_p、broadcast_in_dim_p、dynamic_slice_p、gather_p、scatter_p等; - 归约/窗口:
reduce_sum_p、reduce_max_p、reduce_window_p、select_and_scatter_p等; - 控制流:
cond_p、while_p、scan_p、cumsum_p等; - 并行通信:
psum_p、pmax_p、all_gather_p、all_to_all_p、ppermute_p等; - 线性代数:
dot_general_p、cholesky_p、eigh_p、qr_p、svd_p、lu_p、triangular_solve_p等; - 随机数:
random_bits_p、random_split_p、random_fold_in_p、threefry2x32_p等; - 变换包装:
jit_p、sharding_constraint_p、custom_jvp_call_p、custom_vjp_call_p、remat_p、stop_gradient_p、call_p等。
call_p/closed_call_p值得单独说明:源码中它们是eval_jaxpr_p的别名(core.py),是"把闭式 jaxpr 作为整体调用"的原语,也是JaxprEqn中出现"内嵌调用"时的载体。
四、Effect 效果系统
JAX 用显式效果集合追踪"带副作用"的运算(如 IO、随机数状态、可变引用读写),这直接影响 jit 缓存、控制流合法性、自动微分与 remat 的取舍。
Effect(jax/_src/effects.py):所有效果的基类,本身只是"一种通用副作用"的标记;Effects:Effects = Set[Effect](effects.py),效果集合就是Effect的 Python 集合;no_effects:no_effects: Effects = frozenset()(effects.py),空效果集合,是绝大多数纯运算方程的默认值(Jaxpr构造器默认effects=no_effects);- 配套类型
JaxprInputEffect(effects.py)表示"与某个输入关联的效果":在抽象求值阶段用整数位置指代输入,形成方程时由core.resolve_input_effects(core.py)解析为具体的Var。
此外 effects.py 定义了EffectTypeSet(按类型过滤效果集合的容器)以及一系列全局注册表:ordered_effects、shardable_ordered_effects、lowerable_effects、control_flow_allowed_effects、custom_derivatives_allowed_effects、remat_allowed_effects、partial_eval_kept_effects。例如GenericEffect(core.py)在创建时就被注册进lowerable_effects、control_flow_allowed_effects、custom_derivatives_allowed_effects,因此用def_effectful_abstract_eval2声明效果的原语自动具备被控制流、自定义导数接受的资格。
五、类型抽象、Token 与 jaxtype 判定
5.1array_types与valid_jaxtype
array_types(jax/_src/abstract_arrays.py):array_types = {literals.TypedNdArray, np.ndarray} | numpy_scalar_types,即"JAX 认可的数组/标量 Python 类型集合"。numpy_scalar_types覆盖 int4/int8/.../int64、uint4/.../uint64、complex64/128、bool 及全部浮点标量类型;valid_jaxtype(x) -> bool(core.py):尝试对x求抽象值,若成功且不是字符串 dtype,返回True,否则False。这是快速判定"某 Python 对象能否作为 JAX 值参与计算"的实用工具(字符串数组被显式排除)。
5.2AbstractToken与Token
JAX 用 token 表达"必须在时间上排序"的依赖(如 host callback、IO):
AbstractToken(core.py):token 的抽象值,str_short()显示为Tok,作为切线/余切值是其自身;全局单例abstract_token;Token(core.py):具体 token 对象,内部包裹一个Array缓冲区_buf,用于把数据依赖"线程化"进出计算,提供block_until_ready()。它已注册进pytype_aval_mappings和 canonicalize 处理器,因此可以作为合法 JAX 值参与跟踪。
六、追踪(Tracing)机制相关符号
JAX 的变换(jit/vmap/grad 等)依赖"跟踪器 + trace 栈"机制。jax.extend.core暴露了以下追踪相关符号:
TraceTag(core.py):标识"一组预先存在的 tracers"的标签。源码注释提醒它的相等/哈希实现(所有TraceTag实例互相相等)依赖"函数变换由 tag 参数化、外层函数不可能闭包捕获 trace"这一微妙前提,主要用于缓存键计算;set_current_trace(trace, check_leaks=False)(core.py):上下文管理器,把指定 trace 设为当前 trace,退出时恢复;若check_leaks=True且jax_check_tracer_leaks配置开启,退出时会检查并报告泄漏的 tracers;take_current_trace()(core.py):上下文管理器,返回当前 trace 并临时把当前 trace 置空(用于阻止 trace 逃逸),退出恢复;get_opaque_trace_state(convention=None)(core.py):返回当前 trace 的不透明引用(OpaqueTraceState,基于 weakref 且可按 trace 相等性比较),可用于把"当前追踪状态"放进缓存键而不引入强引用;find_top_trace(_)(core.py):历史遗留函数,等价于取当前 trace,源码标注TODO(douglam): deprecate/delete;nonempty_axis_env_DO_NOT_USE()(core.py):当前轴环境(axis_env.axis_sizes)是否非空,即当前是否处于带命名的 vmap 轴环境内;unsafe_am_i_under_a_jit_DO_NOT_USE()/unsafe_am_i_under_a_vmap_DO_NOT_USE()(core.py):通过检查 trace 栈的字符串表示判断当前是否处于 jit / vmap 变换内(不透明且脆弱,仅作兼容保留,名字已明确告诫勿用);unsafe_get_axis_names_DO_NOT_USE():获取当前轴环境中的命名轴,同样是不稳定 API,仅供兼容旧代码。
七、jaxpr 构造、校验与执行
7.1 构造工具
new_jaxpr_eqn(invars, outvars, primitive, params, effects, source_info=None, ctx=None)(core.py):构造一条JaxprEqn的推荐入口。它会自动补齐source_info与JaxprEqnContext,解析输入相关效果(resolve_input_effects),并在enable_checks开启时断言输入均为Var/Literal、输出均为Var;jaxpr_as_fun(closed_jaxpr)(core.py):把闭式 jaxpr 变成可调用的 Python 函数(柯里化实现),内部在临时关闭debug_nans的情况下调用eval_jaxpr,返回所有输出;gensym:见 2.3 节,用于生成变量;Jaxpr构造器与replace()(core.py):支持直接Jaxpr(constvars, invars, outvars, eqns, effects, debug_info, is_high, consts)显式构造,replace()支持按字段重建;旧式ClosedJaxpr(jaxpr, consts)与replace(jaxpr=..., consts=...)调用形式也仍被兼容。
7.2 校验:check_jaxpr与JaxprTypeError
check_jaxpr(jaxpr)(core.py)是官方提供的 jaxpr 良构性检查器,检查内容包括:
- 被读取的变量必须在此之前被绑定;
- 变量在 jaxpr 全程类型一致;
- 变量类型标注与其绑定表达式兼容。
校验失败时抛出JaxprTypeError(core.py,TypeError子类),并在错误信息中附上出错方程前后各 10 条的格式化 jaxpr 片段以辅助定位。当jax_debug_key_reuse配置开启时,还会额外运行随机数密钥复用检查。JaxprEqn类型检查规则可通过custom_typechecks注册表扩展(如eval_jaxpr_p的闭式调用检查,见 core.py)。
7.3 执行:eval_jaxpr
虽然eval_jaxpr本身未列入本页 autosummary(由jax.extend.core.primitives导入的create_call_primitive等间接使用),但理解 core.py 中的eval_jaxpr(jaxpr, consts, *args)有助于把握整体语义:它把常量和实参写入环境,逐条方程调用eqn.primitive.bind(使用方程的source_info与ctx恢复上下文),多输出原语按序写入多个outvars,最后返回outvars的求值结果。JIT 编译后的 XLA 执行路径与这个解释器共享同一份 IR 语义。
八、其余实用符号
8.1concrete_or_error与InconclusiveDimensionOperation
concrete_or_error(force, val, context="")(core.py):尝试对val求具体值。若val是 Tracer,尝试to_concrete_value();取不到具体值(如被 vmap/jit 的符号维度)则抛出ConcretizationTypeError并携带context信息;force=None时退化为恒等函数。这是"必须拿到编译期常量"场景(如数组长度)的标准工具;InconclusiveDimensionOperation(core.py):jax.errors命名空间下的异常,当"无法对符号维度做结论性计算"时抛出,是形状多态(shape polymorphism)相关代码的哨兵异常。
8.2 vmap 轴映射:mapped_aval/unmapped_aval
mapped_aval(size, axis, aval)(core.py):给定轴大小与轴位置,返回该抽象值"被映射"(增加一个 batch 轴)后的抽象值;通过aval_mapping_handlers注册表按类型分派,未注册类型抛TypeError;unmapped_aval(size, axis, aval, explicit_mesh_axis=None)(core.py):反向操作,移除 batch 轴。它们与 jax/_src/hijax.py 中的HiType(高/低层 IR 类型系统)交互,是 vmap 解释器处理抽象值时的底层支撑。
8.3 自动微分辅助:primal_dtype_to_tangent_dtype
primal_dtype_to_tangent_dtype(primal_dtype)(core.py):返回给定 primal dtype 对应的切线 dtype。规则为:扩展 dtype 走其注册的tangent_dtype规则;非浮点(整数、布尔等)dtype 返回dtypes.float0(零维浮点占位类型,用于"形状正确但梯度恒为零"的整数输入);浮点 dtype 原样返回。这是自定义 VJP/JVP 规则中处理整数参数的标准约定。
8.4DebugInfo
DebugInfo(core.py)实际来自 jax/_src/linear_util.py(DebugInfo = lu.DebugInfo),携带函数名、参数名(arg_names)、结果路径(result_paths)等调试元数据,用于生成报错信息与make_jaxpr的可读输出。Jaxpr构造时会调用debug_info.resolve_result_paths()并(在enable_checks下)校验参数名/结果路径与输入输出数量一致(core.py)。
九、实战:用jax.extend.core定义原语与遍历 IR
综合以上 API,一个典型的"扩展者"工作流如下(示意,基于本仓库 API 形态):
import jax.extend.core as jec from jax.extend.core import primitives as jp # 1) 定义一个新原语:绑定实现与抽象求值 my_p = jec.Primitive("my_op") my_p.def_impl(lambda x: x + 1) # eager 执行 my_p.def_abstract_eval(lambda x: x) # 类型/形状推断(无副作用) # 2) 查看 jaxpr 结构:遍历方程、提取子 jaxpr jaxpr = jax.make_jaxpr(lambda x: x * 2)(jnp.ones(3)).jaxpr for eqn in jaxpr.eqns: print(eqn.primitive.name, eqn.invars, eqn.outvars, eqn.params) for sub in jec.subjaxprs(jaxpr): # 控制流原语内嵌的 jaxpr print(sub) # 3) 校验手写的 jaxpr jec.check_jaxpr(my_jaxpr) # 非法结构会抛 jec.JaxprTypeError # 4) 把闭式 jaxpr 变成可调用函数 fun = jec.jaxpr_as_fun(closed_jaxpr)几点工程建议:
- 变更检测:由于
jax.extend.core无兼容性保证,升级 JAX 后应关注 CHANGELOG.md 中jax.extend相关条目,必要时固定 JAX 版本; - 避免不稳定符号:
unsafe_*_DO_NOT_USE、nonempty_axis_env_DO_NOT_USE等仅用于 JAX 内部兼容,扩展代码不应依赖; - 效果声明:自定义带副作用的原语时,用
def_effectful_abstract_eval返回(out_aval, {effect}),并把自定义Effect类型注册进effects.py中的lowerable_effects/control_flow_allowed_effects等EffectTypeSet,否则控制流与变换可能拒绝处理; - 符号维度:涉及动态形状的抽象求值不要假设维度一定是整数,必要时捕获
InconclusiveDimensionOperation并回退到符号计算。
十、延伸阅读
- 模块总览与设计动机:docs/jep/15856-jex.md、docs/jax.extend.rst
- 核心实现:jaxpr / 原语 / 追踪机制全部集中在 jax/_src/core.py,效果系统见 jax/_src/effects.py,类型映射与
array_types见 jax/_src/abstract_arrays.py - 原语全集索引:jax/extend/core/primitives.py
- 常量的内联/提升规则:docs/internals/constants.md
- 手写解释器与 IR 遍历教程:docs/notebooks/Writing_custom_interpreters_in_Jax.md、docs/601/jaxpr.md
本文所有 API 行为均以当前仓库源码为准。由于jax.extend定位为不稳定二级 API,任何符号在后续版本中都有可能调整,请以本仓库 CHANGELOG.md 与源码为准进行核对。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考