JAXjax.extend.core深入解析Jaxpr 中间表示、Primitive 原语与底层核心机制【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxjax.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.extendjax.extend模块定义见 jax/extend/init.py由 JEP #15856对应仓库文档 docs/jep/15856-jex.md提出目的是把 JAX 的内部组件整理成一个二级 API 库视图供 JAX 生态中的下游库如 Oryx、jax-triton 以及各类自定义变换/编译器前端使用。其关键定位是无兼容性保证jax.extend不遵循公开 API 的兼容性策略不承诺弃用窗口也不承诺跨版本向后兼容破坏性变更通过 CHANGELOG.md 公布区别于jax.experimentalexperimental是新特性的试验场最终可能进入其他模块或被移除而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_arraysarray_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后缀表明它们是出于兼容性保留的内部实现细节新代码不应依赖。二、JaxprJAX 的核心中间表示JaxprJAX Expression是 JAX 变换的核心中间表示。jax.make_jaxpr把 Python 函数跟踪trace后得到的就是一纸 Jaxpr。jax.extend.core提供了一整套用于表示和操作 Jaxpr 的类型。2.1Jaxpr与ClosedJaxprJaxpr类定义在 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 的整体效果集合Effectsdebug_info调试信息DebugInfois_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或Literaloutvars: list[Var]输出变量多输出原语如eigh_p时是多个primitive: Primitive对应的原语params: dict[str, Any]编译期静态参数如dot_general的维度约定effects: Effects本条方程的效果source_info源代码位置与 name stackctx: JaxprEqnContext方程创建时的求值上下文快照compute_type、抽象 mesh、布局模式、xla_metadata 等见 jax/_src/core.py。源码注释特别说明JaxprEqn刻意用带__slots__的普通类而不是 NamedTuple因为构造方程是热路径需要更快的速度。它提供replace()方法用于在改写 IR 时生成部分字段更新的副本。2.3Var、Literal、DropVar与AtomVarcore.pyjaxpr 中的变量节点__slots__只有一个字段aval其抽象值。repr形如Var(id...):f32[3,4]gensymcore.py一个工厂函数gensym lambda: Var用来在构造 jaxpr 时生成全新的变量类型配合Literal语义通常用于中间变量占位Literalcore.py常量节点持有val值与aval抽象值两个字段。只有literalable_types/literalable_scalar_types中登记的类型如 NumPy 标量、TypedNdArray、Python 标量见 jax/_src/abstract_arrays.py才会被嵌入 jaxpr 成为Literal其余常量会被提升为constvarsDropVarcore.pyVar的特例表示该赋值结果永远不会被读取被丢弃的输出美化打印为_Atom类型别名Atom Var | Literalcore.py用于标注方程输入/输出既可以来自变量也可以来自常量的位置。关于常量还有一组配套工具is_literalable(x, for_adFalse)判断值能否作为Literal内联for_adTrue时保留常量以便自动微分is_hoistable(v)判断非标量大常量是否需要提升为函数参数阈值由配置embedded_constants_max_bytes控制jaxpr_const_args(jaxpr)收集 jaxpr 中需要提升的非标量常量。这些细节与 docs/internals/constants.md 中描述的常量机制一一对应。2.4 jaxpr 遍历与子 jaxprjaxprs_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_primitivefinal-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取出当前 traceprev_trace trace_ctx.trace临时将全局 trace 置空再调用bind_with_tracebind_with_tracecore.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_highTrue的子 jaxpr这是 JAX 高/低两层 IRhigh-level / low-level jaxpr机制的入口之一。3.3 规则注册Primitive提供一组def_*便捷方法用于注册各阶段规则方法作用def_impl(impl)注册直接求值实现impl默认抛NotImplementedErrordef_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_evalcore.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 的取舍。Effectjax/_src/effects.py所有效果的基类本身只是一种通用副作用的标记EffectsEffects Set[Effect]effects.py效果集合就是Effect的 Python 集合no_effectsno_effects: Effects frozenset()effects.py空效果集合是绝大多数纯运算方程的默认值Jaxpr构造器默认effectsno_effects配套类型JaxprInputEffecteffects.py表示与某个输入关联的效果在抽象求值阶段用整数位置指代输入形成方程时由core.resolve_input_effectscore.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。例如GenericEffectcore.py在创建时就被注册进lowerable_effects、control_flow_allowed_effects、custom_derivatives_allowed_effects因此用def_effectful_abstract_eval2声明效果的原语自动具备被控制流、自定义导数接受的资格。五、类型抽象、Token 与 jaxtype 判定5.1array_types与valid_jaxtypearray_typesjax/_src/abstract_arrays.pyarray_types {literals.TypedNdArray, np.ndarray} | numpy_scalar_types即JAX 认可的数组/标量 Python 类型集合。numpy_scalar_types覆盖 int4/int8/.../int64、uint4/.../uint64、complex64/128、bool 及全部浮点标量类型valid_jaxtype(x) - boolcore.py尝试对x求抽象值若成功且不是字符串 dtype返回True否则False。这是快速判定某 Python 对象能否作为 JAX 值参与计算的实用工具字符串数组被显式排除。5.2AbstractToken与TokenJAX 用 token 表达必须在时间上排序的依赖如 host callback、IOAbstractTokencore.pytoken 的抽象值str_short()显示为Tok作为切线/余切值是其自身全局单例abstract_tokenTokencore.py具体 token 对象内部包裹一个Array缓冲区_buf用于把数据依赖线程化进出计算提供block_until_ready()。它已注册进pytype_aval_mappings和 canonicalize 处理器因此可以作为合法 JAX 值参与跟踪。六、追踪Tracing机制相关符号JAX 的变换jit/vmap/grad 等依赖跟踪器 trace 栈机制。jax.extend.core暴露了以下追踪相关符号TraceTagcore.py标识一组预先存在的 tracers的标签。源码注释提醒它的相等/哈希实现所有TraceTag实例互相相等依赖函数变换由 tag 参数化、外层函数不可能闭包捕获 trace这一微妙前提主要用于缓存键计算set_current_trace(trace, check_leaksFalse)core.py上下文管理器把指定 trace 设为当前 trace退出时恢复若check_leaksTrue且jax_check_tracer_leaks配置开启退出时会检查并报告泄漏的 tracerstake_current_trace()core.py上下文管理器返回当前 trace 并临时把当前 trace 置空用于阻止 trace 逃逸退出恢复get_opaque_trace_state(conventionNone)core.py返回当前 trace 的不透明引用OpaqueTraceState基于 weakref 且可按 trace 相等性比较可用于把当前追踪状态放进缓存键而不引入强引用find_top_trace(_)core.py历史遗留函数等价于取当前 trace源码标注TODO(douglam): deprecate/deletenonempty_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_infoNone, ctxNone)core.py构造一条JaxprEqn的推荐入口。它会自动补齐source_info与JaxprEqnContext解析输入相关效果resolve_input_effects并在enable_checks开启时断言输入均为Var/Literal、输出均为Varjaxpr_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与JaxprTypeErrorcheck_jaxpr(jaxpr)core.py是官方提供的 jaxpr 良构性检查器检查内容包括被读取的变量必须在此之前被绑定变量在 jaxpr 全程类型一致变量类型标注与其绑定表达式兼容。校验失败时抛出JaxprTypeErrorcore.pyTypeError子类并在错误信息中附上出错方程前后各 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与InconclusiveDimensionOperationconcrete_or_error(force, val, context)core.py尝试对val求具体值。若val是 Tracer尝试to_concrete_value()取不到具体值如被 vmap/jit 的符号维度则抛出ConcretizationTypeError并携带context信息forceNone时退化为恒等函数。这是必须拿到编译期常量场景如数组长度的标准工具InconclusiveDimensionOperationcore.pyjax.errors命名空间下的异常当无法对符号维度做结论性计算时抛出是形状多态shape polymorphism相关代码的哨兵异常。8.2 vmap 轴映射mapped_aval/unmapped_avalmapped_aval(size, axis, aval)core.py给定轴大小与轴位置返回该抽象值被映射增加一个 batch 轴后的抽象值通过aval_mapping_handlers注册表按类型分派未注册类型抛TypeErrorunmapped_aval(size, axis, aval, explicit_mesh_axisNone)core.py反向操作移除 batch 轴。它们与 jax/_src/hijax.py 中的HiType高/低层 IR 类型系统交互是 vmap 解释器处理抽象值时的底层支撑。8.3 自动微分辅助primal_dtype_to_tangent_dtypeprimal_dtype_to_tangent_dtype(primal_dtype)core.py返回给定 primal dtype 对应的切线 dtype。规则为扩展 dtype 走其注册的tangent_dtype规则非浮点整数、布尔等dtype 返回dtypes.float0零维浮点占位类型用于形状正确但梯度恒为零的整数输入浮点 dtype 原样返回。这是自定义 VJP/JVP 规则中处理整数参数的标准约定。8.4DebugInfoDebugInfocore.py实际来自 jax/_src/linear_util.pyDebugInfo 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 PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
