JAX 闭包常量(Closed-over Constants)处理机制深度解析:从 Jaxpr 追踪到 HLO 常量提升(Hoisting)的完整设计
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
导读
本文基于 JAX 官方内部文档 docs/internals/constants.md,深入剖析 JAX 如何追踪与降级(lowering)那些在函数追踪期被"无意中"捕获的非标量常量(closed-over constants)。你将了解到这些常量在Jaxpr中如何以core.Literal表示、在 lowering 阶段如何被提升(hoist)为额外的函数参数(const_args)以避免内联进 HLO、以及新的简化实现(由JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS开启)与旧的ClosedJaxpr实现之间的差异。读完本文,你将能够理解jax.jit编译管线中常量处理的关键路径,并学会使用JAX_CAPTURED_CONSTANTS_WARN_BYTES等配置来诊断闭包常量带来的性能隐患。
什么是闭包常量(Closed-over Constants)
在 JAX 中,闭包常量是指在对一个函数进行追踪(tracing)时遇到的、不依赖该函数任何参数的非标量数组。JAX 的jax.numpy和lax等操作是"stage out"的(即被记录进计算图而不是立即执行),因此它们不会产生闭包常量;而原生的 NumPy 操作或预先构造好的jax.Array则会。
文档给出了一个非常直观的例子:
import numpy as np from jax import jit from jax import numpy as jnp a_jax_array = jnp.ones((16,), dtype=np.float32) @jit def f(x): return x + a_jax_array + np.full((16,), 42.) + jnp.full((16,), 142.)在这个例子中,a_jax_array(预先构造的jax.Array)和np.full((16,), 42.)(NumPy 原生的ndarray)都是闭包常量;而jnp.full((16,), 142.)是 JAX 操作,在追踪时被记录为计算图节点,不是闭包常量。
闭包常量为何值得警惕
闭包常量最容易在不知不觉中被引入。典型场景包括:
- 在
jitted函数体外预先计算好的权重矩阵、掩码(mask)或索引表被函数体直接引用; - 在函数体内部直接调用 NumPy 函数(如
np.ones、np.arange),这些结果会在追踪时被物化为常量嵌入计算图; - 从数据加载流程中读入的、形状与函数参数无关的辅助数据。
当这些常量较大时,它们会被内联进 HLO 代码,导致后续一系列问题(详见 Lowering 阶段的取舍)。
使用 JAX_CAPTURED_CONSTANTS_WARN_BYTES 诊断意外捕获
文档指出,可以设置环境变量JAX_CAPTURED_CONSTANTS_WARN_BYTES为任意非负值,从而在函数 lowering 期间记录(警告)所有不小于该字节数的闭包常量,帮助你发现意外捕获。
从 jax/_src/config.py 的源码可以看到该配置的真实定义:
captured_constants_warn_bytes = int_state( name='jax_captured_constants_warn_bytes', default=2 * 10 ** 9, help=('The number of bytes of parameters that may be captured as constants ' 'before a warning is issued. Defaults to approximately 2GB. ' 'Set to -1 to disable issuing a warning.' ) )关键信息:
| 配置项 | 默认值 | 说明 |
|---|---|---|
jax_captured_constants_warn_bytes | 2 * 10 ** 9(约 2GB) | 捕获常量总字节数超过该阈值时发出警告;设为-1可彻底禁用警告 |
jax_captured_constants_report_frames | 0 | 报告中为每个捕获常量显示的调用栈帧数;-1打印完整帧,0禁用报告。注意:仅当捕获常量总字节数超过警告阈值时才生成报告(生成报告开销较大) |
在 jax/_src/interpreters/mlir.py 中,check_jaxpr_constants与log_closed_over_constant实现了该警告逻辑:当closed_jaxpr.consts的nbytes总和超过阈值时,warnings.warn会提示"大量常量在 lowering 期间被捕获(共 N 字节)",并建议要么确认这是有意的,要么通过JAX_CAPTURED_CONSTANTS_WARN_BYTES=-1关闭警告;如需定位捕获位置,可设置JAX_CAPTURED_CONSTANTS_REPORT_FRAMES=-1获取栈帧报告。
新实现概览:JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS
文档强调:以下描述的是未来(文档写作时点为 2026 年 4 月)的常量内部实现细节。它还不是当前默认实现,需要通过环境变量显式开启:
JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS=True源码 jax/_src/config.py 中对这个开关的定义佐证了这一点:
use_simplified_jaxpr_constants = bool_state( name='jax_use_simplified_jaxpr_constants', default=False, help=('Enable a simplification of the handling of closed-over constants ' 'in Jaxpr. The value `True` enables the new behavior. ' 'This flag will exist only briefly, while we transition ' 'users. See https://docs.jax.dev/en/latest/internals/constants.html.' 'DO NOT RELY ON THIS FLAG.'), include_in_jit_key=True, include_in_trace_context=True)注意两点:
- 该 flag 的
include_in_jit_key=True、include_in_trace_context=True,意味着它会参与 jit 缓存键与追踪上下文的构成——不同取值下编译出的可执行文件不能混用缓存; - 源码注释明确警告"DO NOT RELY ON THIS FLAG",这是一个过渡期标志,不应在用户代码中长期依赖。旧实现的细节及其缺陷见 Previous implementation(旧实现)。
Tracing 阶段:core.Literal 与 is_literalable
当 JAX 追踪遇到一个常量——无论它是某个 JAX primitive(算子)的参数,还是函数的返回值——它会被表示为core.Literal,并随使用它的 primitive 一起内嵌在Jaxpr中。
决定哪些常量会被转换为core.Literal的函数是core.is_literalable。根据 jax/_src/core.py 的实现:
- 所有标量常量都会被转换为
core.Literal(literalable_scalar_types走快速路径直接返回True); - 非标量的
np.ndarray与jax.Array也会被转换为core.Literal; - 当
use_simplified_jaxpr_constants开启时,jax.Array(ArrayImpl)在非for_ad场景下也会字面化(do_lit_array = not for_ad),这是为了在自动微分(AD)下保留常量; - 其余类型(例如自定义 Python 对象)则落入选集
literalable_types,仅在满足条件时字面化,否则以constvars(闭包变量)形式出现在Jaxpr上。
同时,core.is_hoistable(jax/_src/core.py)判断一个Literal是否需要被提升为参数:
def is_hoistable(v: Literal) -> bool: return (np.ndim(v.val) > 0 and getattr(v.val, "nbytes", 4) > config.embedded_constants_max_bytes.value)即:非标量且字节数超过embedded_constants_max_bytes的常量才值得提升;小常量会被直接内嵌(见下文)。
Lowering 阶段:常量提升(Hoisting)为 const_args
为什么不直接内联 stablehlo.constant
理论上,lowering 到 HLO 时,最简单的方式是为每个core.Literal直接发射一个stablehlo.constant操作。但文档明确列出了这样做的一系列弊端:
- 主机内存压力与分片丢失:如果常量是
jax.Array(如例子中的a_jax_array),lowering 期间会把它从设备拉回主机,可执行模块执行时再重新物化到设备上。这会显著增加主机内存占用(有时是数量级的增长);更进一步,如果常量在多个设备上做了分片(sharding),这种分片信息在拉回-重新物化的过程中会丢失。 - HLO 膨胀与编译变慢:大常量(尤其被多次复用的同一个常量)会显著增大 HLO 体积;XLA 编译器还会尝试对它们做常量折叠(constant-folding),引发告警并拖慢编译。
- 数值差异风险:实测中 XLA 的常量折叠有时会产生与编译后代码略有不同的数值结果。
jaxpr_const_args:扫描并去重常量
文档指出,lowering 期间使用core.jaxpr_const_args来扫描一个Jaxpr,返回其中包含的常量列表,按id去重(uniquified)。该函数对每个Jaxpr及其子Jaxpr调用结果会被记忆化(memoized)。
看 jax/_src/core.py 的真实实现:
@partial(weakref_lru_cache, trace_context_in_key=False) def jaxpr_const_args(jaxpr: Jaxpr) -> list[tuple[ArrayLike, AbstractValue]]: # The non-scalar constants in core.Literal, in the entire Jaxpr, # uniquified by id. These will be hoisted as const arguments to the functions # in which they appear. if not config.use_simplified_jaxpr_constants.value: return [] consts_by_id: dict[int, tuple[ArrayLike, AbstractValue]] = {} for v in jaxpr.outvars: if type(v) is Literal and is_hoistable(v): consts_by_id[id(v)] = (v.val, v.aval) for eqn in jaxpr.eqns: for v in eqn.invars: if type(v) is Literal and is_hoistable(v): consts_by_id[id(v)] = (v.val, v.aval) consts_by_id.update({id(v_aval[0]): v_aval for v_aval in eqn_params_const_args(eqn.params)}) return list(consts_by_id.values())实现要点:
- 通过
weakref_lru_cache记忆化(同时以id哈希为基础,因此同一常量不会重复扫描); - 只收集
is_hoistable(非标量、字节数超过embedded_constants_max_bytes)的Literal; - 遍历
outvars与每个方程的invars,同时通过eqn_params_const_args递归收集方程参数中嵌套Jaxpr(子函数)的常量; - 在
use_simplified_jaxpr_constants=False(默认)时直接返回空列表,即旧行为不受影响。
const_args 的参数排布与 const_lowering 映射
所有被降级的 HLO 函数都会为Jaxpr中出现的每个唯一常量多接收一个额外参数。这些参数称为const_args,其排布位置是:
维度变量参数(dimension variable args)之后 → token 参数之后 → 实际数组参数(array arguments)之前
lowering 期间维护一个映射:
const_lowering: dict[int, mlir.IrValues]该映射以常量的id为键,值为对应的 HLO 值,被存放在mlir.LoweringRuleContext中。mlir.ir_constant在遇到常量时会优先复用const_lowering中已有的 lowering,而不是重新发射stablehlo.constant(见 jax/_src/interpreters/mlir.py,其中_ir_constant会在const_lowering命中时直接复用既有值)。
小常量例外:embedded_constants_max_bytes
存在一个例外:尺寸不超过config.embedded_constants_max_bytes的小常量不会被提升为参数,而是直接内嵌(embed)进生成的 HLO 与可执行文件中。该配置定义于 jax/_src/config.py:
embedded_constants_max_bytes = int_state( name='jax_embedded_constants_max_bytes', default=32, help=('Maximum size in bytes of a constant that is allowed to be ' 'embedded in the lowered HLO. Constants larger than this ' 'are hoisted as additional arguments to the executable. ' 'See https://docs.jax.dev/en/latest/internals/constants.html.'), include_in_jit_key=True, include_in_trace_context=True)默认值为32 字节。也就是说,小于等于 32 字节的非标量常量(以及所有标量常量)仍以内联stablehlo.constant形式存在,方便 XLA 做常量折叠;大于 32 字节的常量才被提升为const_args。与use_simplified_jaxpr_constants一样,它同样参与 jit 缓存键与追踪上下文。
内部函数(inner function)的 lowering
当 lowering 一个 HLO 内部函数(非main函数)时,会再次调用core.jaxpr_const_args获取对应Jaxpr中实际的常量。这些常量预期已经包含在外层函数的const_lowering中;内部函数会获得自己更小的一组const_args和自己的const_lowering映射,用于 lowering 其函数体。文档举例mlir.lower_jaxpr_as_fun就是发生此类逻辑的一处。
而mlir.jaxpr_subcomp(jax/_src/interpreters/mlir.py)不会创建新的 HLO 函数,而是在当前函数内创建一个 block,并复用外层函数的const_lowering。
仍会出现的 stablehlo.constant
文档特别说明,即便在新实现下,降级代码中依然会存在stablehlo.constant,出现在以下四种场景:
- 标量常量:希望将这些常量暴露给 XLA 做常量折叠;
- 小常量:尺寸不超过
embedded_constants_max_bytes(默认 32 字节)的常量,如上文所述直接内嵌; - lowering 期间新产生的常量:常量未出现在被追踪的程序中(因此不在
Jaxpr里)。例如某些 PRNG(随机数)函数的 lowering 就自带了常量; - 导出(export)场景:目前导出时不提升常量参数,因为导出序列化尚不支持数组序列化。这是通过
mlir.LoweringParameters.hoist_constants_as_args参数控制的(其默认值与use_simplified_jaxpr_constants一致,见 jax/_src/interpreters/mlir.py)。
avals、shardings 与 layouts 的高层计算
还有一个实现细节:部分内部 lowering 函数需要用到参数 avals,有时还需要参数的 shardings 与 layouts。而且包括const_args在内的所有参数的 avals、shardings、layouts 在 lowering 之后也仍然会被使用。因此,比较方便的做法是在调用栈的较上层一次性算好,例如在pxla.lower_sharding_computations中计算并向下传递。
具体来说,mlir.lower_jaxpr_to_module、pjit._pjit_cached_lower_jaxpr_to_fun、mlir.lower_jaxpr_to_fun这些函数都接收in_avals、in_shardings、in_layouts(这些列表同时包含const_args的 avals 与常规参数的 avals,后者对应Jaxpr.invars),此外还接收一个num_const_args参数用于区分常量参数与常规参数。
编译与执行:const_args 如何传入可执行文件
lowering 出的 MLIR 模块包含 const_args 对应的参数,因此编译后的可执行文件在被调用时也必须传入 const_args。这里的关键设计问题是:在哪个位置把 const_args 拼接到调用参数前面。
文档给出了一个示例,强调第二次调用应命中 C++ jit 缓存而不执行任何 Python 代码:
const = jnp.array([42.]) f = jax.jit(lambda: const) f() f()这意味着const必须以某种方式在 C++ 侧传给可执行文件(因此被存储在pxla.MeshExecutableFastpathData中)。相应地,C++ 缓存未命中函数(例如pjit._cpp_pjit.cache_miss,或pxla.MeshExecutable.create_cpp_call中的aot_cache_miss)不接收 const_args 作为参数,而是由这些缓存未命中函数负责自行前置拼接(prepend)const_args。
关于 C++ 快速路径(fast path)的支持情况:
- 从jaxlib 0.7.1开始,C++ 快速路径支持 const_args;
- 在更早的版本中,只要存在 const_args,快速路径就会被禁用(回退到较慢的 Python 路径)。
const_args 在 stage 对象中的存放
为实现上述方案,const_args被保存在以下对象中:
stages.Loweringstages.Loweredstages.CompiledCallParamspxla.MeshExecutable
注意:在stages.Compiled中,in_avals等字段不包含const_args(即Compiled对外呈现的接口不含常量参数)。
序列化(编译缓存)与 const_args
一个有趣的推论是:当序列化可执行文件(例如用于编译缓存)时,无需序列化闭包常量本身——可执行文件本身不包含这些常量,它只是需要接收它们作为 const_args。因此,反序列化缓存的可执行文件的一方,必须自行提供 const_args。这要求编译缓存的消费者在缓存命中时仍能拿到与编译时一致的闭包常量。
AOT 模式与 x64 的一致性要求
在 AOT(预先编译)模式下,lowering 与执行可能使用不同的jax_enable_x64配置值。文档给出约束:如果常量是 64 位ndarray,那么 lowering 与执行必须使用相同的jax_enable_x64值,否则常量解释会不一致,可能导致错误结果或崩溃。
Previous implementation(旧实现)与缺陷
当JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS=False时(即文档写作时点的默认行为),采用的是 2025 年 7 月的旧方案:
当 JAX 将函数追踪成Jaxpr时,会把闭包值收集进一个常量集合,并给Jaxpr加上一组对应的constvars(真正的函数参数由invars表示)。大多数追踪函数(如trace_to_jaxpr_dynamic)会同时返回Jaxpr和这些常量。代码中大量使用core.ClosedJaxpr类,它封装了一个Jaxpr以及与其constvars对应的consts。
文档明确列出了ClosedJaxpr方案的若干问题:
- 内联问题:
ClosedJaxpr中consts的 lowering 会直接产生内联的stablehlo.constant,即前文描述的各种弊端(主机内存、HLO 膨胀、常量折叠数值差异、分片丢失)。 - 类型混淆:
Jaxpr与ClosedJaxpr在 JAX 中无处不在,且常被笼统地命名为jaxpr,难以区分当前拿到的是哪一种。虽然已开始添加类型声明,但部分代码仍用isinstance条件分支同时兼容两者。 - 缓存键与记忆化困难:
Jaxpr和ClosedJaxpr有时被用作缓存键,且按id哈希,因此希望记忆化它们的构造。例如pe.closed_jaxpr(位于 jax/_src/interpreters/partial_eval.py)记忆化了ClosedJaxpr的构造,但仅在consts为空时——因为有时常量不可哈希。 - lowering 覆盖不全:处理
ClosedJaxpr中的常量需要额外小心。例如 Mosaic lowering 中尚有未实现非空常量ClosedJaxpr处理的地方(见 jax/_src/pallas/mosaic/lowering.py 附近的相关逻辑)。 - 变换中的额外输入:将闭包常量转成输入后,在各变换(transformations)中需要小心处理这些辅助输入(auxiliary inputs)的传递。
这些缺陷正是新实现(简化 Jaxpr 常量)要解决的问题:把常量显式表示为core.Literal、统一通过jaxpr_const_args去重扫描并按需提升为const_args,从而避免内联stablehlo.constant的各种问题。
实践建议与总结
综合文档与源码,针对闭包常量可以给出如下实践要点:
- 诊断先行:在开发阶段设置
JAX_CAPTURED_CONSTANTS_WARN_BYTES(如JAX_CAPTURED_CONSTANTS_WARN_BYTES=1048576表示 1MB)观察是否有非预期的大常量被捕获;配合JAX_CAPTURED_CONSTANTS_REPORT_FRAMES=-1获取捕获位置的调用栈报告。不需要时用-1关闭警告,避免每次 lowering 都产生告警。 - 理解参数排布:在新实现下,
const_args位于维度变量参数与 token 参数之后、数组参数之前;所有参数(含 const_args)的 avals/shardings/layouts 由调用栈上层统一计算并向下传递。 - 缓存与序列化的约定:C++ jit 缓存命中要求常量以
const_args形式在 C++ 侧传递(jaxlib ≥ 0.7.1);编译缓存反序列化时不包含常量本身,缓存消费者必须自己提供 const_args。 - x64 一致性:AOT 场景下若常量是 64 位
ndarray,必须保证 lowering 与执行使用相同的jax_enable_x64。 - 过渡期标志:
JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS与jax_embedded_constants_max_bytes(默认 32 字节)都是过渡性配置,且参与 jit 缓存键与追踪上下文,不应在用户代码中长期依赖,应关注 JAX 版本演进以迁移到默认行为。
本文的所有关键结论均可在仓库源码中得到印证:核心逻辑 core.py、配置定义 config.py、lowering 实现 mlir.py 与 partial_eval.py。建议读者在阅读本文后,结合上述源码文件与 constants.md 原文,进一步追踪jax.jit从追踪到执行的完整常量处理链路。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考