尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

JAX 内部实现剖析:closed-over 常量的捕获、提权与降级(HLO Lowering)全流程

发布时间:2026/9/10 9:35:26

资讯中心
01
ARTICLE

JAX 内部实现剖析:closed-over 常量的捕获、提权与降级(HLO Lowering)全流程

JAX 内部实现剖析:closed-over 常量的捕获、提权与降级(HLO Lowering)全流程
JAX 内部实现剖析closed-over 常量的捕获、提权与降级HLO Lowering全流程【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax本篇技术指南基于 JAX 官方内部实现笔记 docs/internals/constants.md 展开系统讲解 JAX 在追踪tracing与降级lowering阶段如何处理被函数意外闭包捕获的非标量常量数组从core.Literal的判定、jaxpr_const_args的去重收集到以const_args形式作为额外参数传入 HLO 函数、再交由 C 快速路径缓存与执行的完整链路。读完本文你将理解JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS、JAX_CAPTURED_CONSTANTS_WARN_BYTES、jax_embedded_constants_max_bytes等配置的真实作用并掌握如何诊断与规避大常量内联进 HLO 引发的编译变慢、数值差异与主机内存膨胀问题。什么是 closed-over constants闭包捕获常量JAX 的转换机制jit、vmap、grad等基于追踪tracing当 Python 函数被追踪时其入参被替换为抽象的追踪值而函数体访问到的不依赖任何函数参数的、非标量数组就被称为closed-over constants闭包捕获常量。原始文档给出了一个典型示例import numpy as np from jax import jit from jax import numpy as jnp a_jax_array jnp.ones((16,), dtypenp.float32) jit def f(x): return x a_jax_array np.full((16,), 42.) jnp.full((16,), 142.)在这个例子中a_jax_array一个jax.Array是 closed-over constantnp.full((16,), 42.)一个 NumPyndarray同样是 closed-over constant而jnp.full((16,), 142.)不是closed-over constant——因为jax.numpy与lax下的操作会被staged out延迟为计算图的一部分在追踪时就被提升为 JAX 程序内的计算而不是被当作外部捕获的常量。文档中随后统称这类常量简称为constants。为什么会无意中捕获常量闭包捕获常量很容易在不知不觉中出现例如在jit装饰的函数体内直接引用外层作用域的 NumPy 数组、模型权重缓存、或其它 JAX 数组。为此JAX 提供了诊断开关设置环境变量JAX_CAPTURED_CONSTANTS_WARN_BYTES为任意非负值后在函数降级lowering期间凡是尺寸达到该阈值的 closed-over constants 都会被记录并输出警告。底层实现位于 jax/_src/interpreters/mlir.py 的check_jaxpr_constants与log_closed_over_constant当捕获常量总字节数超过阈值时警告会列出最大的几个常量及其捕获位置调用帧信息并提示如果这是有意为之可通过设置JAX_CAPTURED_CONSTANTS_WARN_BYTES-1关闭该警告若想看到常量具体在哪些等式eqn中被使用还可配合JAX_CAPTURED_CONSTANTS_REPORT_FRAMES设置报告帧数。新实现总览JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS原始文档明确说明下文描述的是未来的内部实现细节。截至文档撰写时间2026 年 4 月该实现还不是默认实现需要通过环境变量开启export JAX_USE_SIMPLIFIED_JAXPR_CONSTANTSTrue该配置项定义于 jax/_src/config.py其帮助文本明确指出这是一个过渡性 flagDO NOT RELY ON THIS FLAG不要依赖该 flag用于在切换期间帮助用户迁移。它同时被纳入include_in_jit_keyTrue与include_in_trace_contextTrue意味着该配置的取值会参与 JIT 缓存键与追踪上下文的构造。旧实现JAX_USE_SIMPLIFIED_JAXPR_CONSTANTSFalse会把常量内联为stablehlo.constant带来一系列问题详见文末旧实现及其缺陷一节新实现的核心思路是将常量作为额外的函数参数const_args提权hoist到 HLO 函数签名中避免大常量被直接内联进 HLO。追踪阶段core.Literal 与 is_literalable在追踪期间当一个常量出现在某个 JAX 原语primitive的参数位置或函数返回值位置时它会被表示为一个core.Literal对象并与使用它的原语一起被嵌入Jaxpr中。core.Literal的定义位于 jax/_src/core.py仅包含val与aval两个字段并通过__slots__优化内存占用其hash属性尝试对常量值取哈希失败时退化为(item, dtype)元组哈希这为后续基于id的去重与缓存提供了基础。决定哪些常量会被转换为core.Literal的判定函数是core.is_literalablejax/_src/core.py其判定逻辑如下标量常量全部可转为Literal代码为标量类型维护了literalable_scalar_types集合并走快速路径直接返回True避免np.ndarray转换开销非标量常量np.ndarray与jax.Array也在可转为Literal的类型之列literalable_types集合for_ad参数在自动微分AD变换下jax.Array默认不再转为Literaldo_lit_array not for_ad这是为了在 AD 过程中保留常量语义其它类型的常量则不会成为Literal最终会以constvars的形式出现在Jaxpr中见旧实现一节。literalable_types与literalable_scalar_types是模块级set由各后端在初始化时向其中注册允许的常量类型jax/_src/core.py。降级阶段从 Jaxpr 到 HLO 的常量提权为什么不能简单地为 Literal 发射 stablehlo.constant如果降级到 HLO 时对每个core.Literal直接发射一条stablehlo.constant会带来若干严重缺点这也是新实现要解决的问题主机内存膨胀如果常量是jax.Array例如上文示例中的a_jax_array降级期间需要把它从设备拉回主机host随后在模块执行时又要在设备上重新物化。这会在主机侧占用内存且当常量较大或较多时会sometimes dramatically地增加主机内存占用更严重的是若该常量原本按多个设备分片sharded分片信息会在这一设备→主机→设备的往返中丢失HLO 体积膨胀大常量会增大 HLO 的体积尤其是同一常量被多次使用时会被重复内联XLA 编译器还会尝试对这些常量做 constant-folding导致编译警告与编译变慢数值差异风险官方观察到 XLA 的 constant-folding 有时会产生与编译后代码略有不同的数值结果。这一点可参考 JAX 仓库的 issueLarge closed-over constants are inlined in the HLO code #29684。const_args把常量提升为函数参数新实现的降级策略如下对应 jax/_src/core.py 的core.jaxpr_const_args调用core.jaxpr_const_args(jaxpr)扫描整个Jaxpr返回其中出现的非标量常量列表并按id去重该函数通过weakref_lru_cache装饰器对每个Jaxpr及子Jaxpr的调用结果进行记忆化memoize避免重复扫描去重遍历的对象包括Jaxpr的outvars输出中的Literal与每个等式eqn.invars输入中的Literal且只保留is_hoistable(v)为真的常量——is_hoistable要求常量维度大于 0 且字节数超过config.embedded_constants_max_bytesjax/_src/core.py对于嵌套在等式参数中的子Jaxpr通过eqn_params_const_args递归收集其中的常量。所有被降级的 HLO 函数都会为每个唯一常量多接收一个额外参数这些参数即const_args。在函数签名中的位置约定是在维度变量参数之后、token 参数之后、真正的数组参数对应Jaxpr.invars之前。降级期间维护一个从常量id到 HLO 值的映射const_lowering: dict[int, mlir.IrValues]存储于mlir.LoweringRuleContext中并由mlir.ir_constant使用jax/_src/interpreters/mlir.py当降级过程再次遇到某个常量时直接复用const_lowering中已有的降级结果而不再发射新的stablehlo.constant。小常量的例外embedded_constants_max_bytes对于尺寸不超过config.embedded_constants_max_bytes的小常量会有一个例外处理它们不会被提权为额外参数而是直接内联进生成的 HLO 与可执行文件中。该配置定义于 jax/_src/config.py默认值为32字节同样被纳入 JIT 缓存键与追踪上下文。这意味着低于 32 字节的小常量例如小型的形状/布局辅助数组仍然以内联常量形式存在享受 XLA constant-folding 的优化而大常量则走提权路径。内层函数的降级当降级一个 HLO 内层函数非main函数时会再次调用core.jaxpr_const_args获取对应Jaxpr中的实际常量这些常量预期已包含在外层已建立的const_lowering中内层函数会获得自己更小的const_args集合与独立的const_lowering映射用于降级函数体。例如 jax/_src/interpreters/mlir.py 的mlir.lower_jaxpr_as_fun就是此类场景之一。而mlir.jaxpr_subcompjax/_src/interpreters/mlir.py不会创建新的 HLO 函数而是在当前函数内创建一个 block因此直接复用外层函数的const_lowering无需重复处理常量参数。哪些场景下 stablehlo.constant 依然存在需要说明的是即使开启新实现降级后的代码中仍会出现stablehlo.constant具体包括四种情形标量常量希望它们对 XLA 可见以进行 constant-folding小常量尺寸不超过config.embedded_constants_max_bytes如上所述降级期产生的常量这些常量没有出现在被追踪的程序Jaxpr中因此不在const_lowering里例如某些 PRNG 函数的降级实现会引入常量导出export场景当前 export 序列化不支持数组序列化因此在导出时不会提权常量参数而是通过mlir.LoweringParameters.hoist_constants_as_args参数来控制该行为该参数默认值取自config.use_simplified_jaxpr_constants.value见 jax/_src/interpreters/mlir.py。avals、shardings 与 layouts 的传递常量提权带来一个额外的复杂性部分内部降级函数需要同时拿到参数包括const_args的 avals、shardings 与 layouts并且这些信息在降级之后执行阶段仍然需要。因此在调用栈较高的位置例如pxla.lower_sharding_computations统一计算出这些信息并向下传递是更便捷的做法。例如 jax/_src/interpreters/mlir.py 的mlir.lower_jaxpr_to_module、pjit._pjit_cached_lower_jaxpr_to_fun与mlir.lower_jaxpr_to_fun都接收同时涵盖const_args与常规参数对应Jaxpr.invars的in_avals、in_shardings、in_layouts以及一个显式的num_const_args参数。lower_jaxpr_to_module的 docstring 也明确注明The inputs already account for the constant arguments.输入已经计入常量参数。编译与执行const_args 如何与 JIT 缓存协同降级后的 MLIR 模块包含了const_args对应的参数因此编译出的可执行程序在执行时必须被传入这些常量参数。关键在于选择正确的前置prepend位置。原始文档给出了一个精妙的测试用例const jnp.array([42.]) f jax.jit(lambda: const) f() f()第二次调用f()应当直接命中 C jit 缓存不执行任何 Python 代码。这就意味着const必须在 C 层被传给可执行程序因此被存放在pxla.MeshExecutableFastpathData中并且 C 缓存未命中cache miss回调函数——例如pjit._cpp_pjit.cache_miss、或pxla.MeshExecutable.create_cpp_call中的aot_cache_miss——不会把const_args作为入参相反这些 cache-miss 函数需要自行前置拼接const_args才能保证第二次调用走全缓存路径。文档在此处留下一句 TODO计划由 yashk2810 补充一份关于 jit 缓存工作机制的描述本文不再展开。关于版本兼容性有一个重要事实C 快速路径从 jaxlib 0.7.1 开始支持const_args在更早的版本中只要存在const_args快速路径就会被禁用回退到较慢的 Python 路径。为实现这一方案const_args被贯穿保存在多个阶段对象中stages.Lowering降级阶段stages.Lowered已降级阶段stages.CompiledCallParams编译调用参数pxla.MeshExecutable网格可执行程序而在stages.Compiled中in_avals等字段不包含const_args——即编译完成后的对象对外呈现的是纯业务参数签名。序列化与编译缓存一个有趣且重要的推论是序列化可执行程序例如写入编译缓存时并不需要序列化 closed-over 常量本身。可执行程序内部并不包含这些常量它只是要求调用方传入const_args因此反序列化缓存可执行程序的一方必须在调用时自行提供对应的const_args。这一设计使得编译缓存的序列化格式与常量解耦。AOT 模式与 x64 的一致性约束在 AOTahead-of-time模式下降级与执行可能使用不同的jax_enable_x64配置值。如果闭包捕获的常量是 64 位ndarray那么降级与执行必须使用相同的jax_enable_x64取值否则会导致 32/64 位语义不一致、常量降级错误。旧实现及其缺陷ClosedJaxpr在不设置JAX_USE_SIMPLIFIED_JAXPR_CONSTANTSTrue时即False截至文档撰写时间 2025 年 7 月仍是默认JAX 采用旧的常量处理方式追踪函数到Jaxpr时JAX 把闭包捕获的值收集为一个常量集合并在Jaxpr上添加对应的constvars真正的函数参数则用invars表示大多数追踪函数例如trace_to_jaxpr_dynamic会同时返回Jaxpr与常量集合代码中广泛使用core.ClosedJaxpr类它封装了一个Jaxpr及其consts对应Jaxpr.constvars。ClosedJaxpr存在以下问题内联常数ClosedJaxpr中consts的降级结果是内联的stablehlo.constant即前文所述全部缺点的根源主机内存膨胀、HLO 体积膨胀、constant-folding 数值差异、分片丢失类型混淆Jaxpr与ClosedJaxpr在 JAX 中被普遍使用且常以通用的jaxpr名字出现难以区分当前拿到的是哪一种虽然已经逐步引入类型声明但有些代码仍依赖isinstance条件分支同时兼容两者缓存与记忆化困难Jaxpr与ClosedJaxpr有时作为缓存键使用、且按id哈希因此希望对其构造进行记忆化。例如pe.closed_jaxprjax/_src/interpreters/partial_eval.py 中的函数对ClosedJaxpr的构造做了记忆化但仅在consts为空时生效——因为有时consts本身不可哈希下游处理缺失处理ClosedJaxpr的常量需要额外小心例如 Mosaic 降级jax/_src/pallas/mosaic/lowering.py中仍有未实现非空常量ClosedJaxpr处理的地方变换期负担把闭包常量转成输入后在各类变换如 AD、vmap中必须小心处理这些辅助输入增加了变换实现的复杂度。新实现通过常量直接以Literal嵌入Jaxpr、降级时按id去重并提权为const_args的方式从根上绕开了上述问题Jaxpr自身自包含不再依赖外部的consts列表、jaxpr_const_args可安全记忆化、大常量不再内联进 HLO。配置项速查表配置 / 环境变量类型默认值作用JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS布尔False启用简化的Jaxpr常量处理常量提权为const_args过渡性 flag勿长期依赖JAX_CAPTURED_CONSTANTS_WARN_BYTES整数未设置非负值时在降级期间记录尺寸达到该阈值的 closed-over constants设为-1关闭JAX_CAPTURED_CONSTANTS_REPORT_FRAMES整数未设置配合上项输出常量捕获位置的调用帧数量-1表示全部报告jax_embedded_constants_max_bytes整数32允许内联进 HLO 的常量最大字节数超过该值的常量被提权为额外参数其中前两项对应源码中的captured_constants_warn_bytes与use_simplified_jaxpr_constants配置jax/_src/config.py后一项即embedded_constants_max_bytes所有与常量相关的配置均被纳入 JIT 缓存键与追踪上下文修改它们会改变缓存行为。结语closed-over constants 的处理是 JAX 追踪/降级管线中最容易被忽视、却又影响深远的一环从core.is_literalable决定哪些常量进入Jaxpr到core.jaxpr_const_args按id去重收集、const_lowering复用降级结果、const_args以固定位置插入 HLO 函数签名再到 C 快速路径缓存与序列化时的不携带常量设计新实现将大常量从 HLO 中彻底剥离换来了更小的编译单元、更快的编译速度与更稳定一致的数值行为。若你在实践中观察到大常量导致 HLO 膨胀或constant-folding 数值异常可先用JAX_CAPTURED_CONSTANTS_WARN_BYTES定位捕获源头再考虑开启JAX_USE_SIMPLIFIED_JAXPR_CONSTANTSTrue验证新路径的效果在正式项目中使用时请以当前发行版的实际默认行为为准并留意该过渡 flag 的未来走向。想要深入源码推荐从以下位置继续阅读核心判定与收集jax/_src/core.pyLiteral、is_literalable、is_hoistable、jaxpr_const_args降级与内联jax/_src/interpreters/mlir.pyir_constant/ir_constants与const_lowering顶层降级入口jax/_src/interpreters/mlir.pylower_jaxpr_to_module捕获常量诊断jax/_src/interpreters/mlir.pycheck_jaxpr_constants/log_closed_over_constant相关配置定义jax/_src/config.py【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。