處理機(jī)制深度解析:從 Jaxpr 追蹤到 HLO 常量提升(Hoisting)的完整設(shè)計(jì))
JAX 閉包常量Closed-over Constants處理機(jī)制深度解析從 Jaxpr 追蹤到 HLO 常量提升Hoisting的完整設(shè)計(jì)【免費(fèi)下載鏈接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more項(xiàng)目地址: https://gitcode.com/GitHub_Trending/ja/jax導(dǎo)讀本文基于 JAX 官方內(nèi)部文檔 docs/internals/constants.md深入剖析 JAX 如何追蹤與降級lowering那些在函數(shù)追蹤期被無意中捕獲的非標(biāo)量常量closed-over constants。你將了解到這些常量在Jaxpr中如何以core.Literal表示、在 lowering 階段如何被提升hoist為額外的函數(shù)參數(shù)const_args以避免內(nèi)聯(lián)進(jìn) HLO、以及新的簡化實(shí)現(xiàn)由JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS開啟與舊的ClosedJaxpr實(shí)現(xiàn)之間的差異。讀完本文你將能夠理解jax.jit編譯管線中常量處理的關(guān)鍵路徑并學(xué)會使用JAX_CAPTURED_CONSTANTS_WARN_BYTES等配置來診斷閉包常量帶來的性能隱患。什么是閉包常量Closed-over Constants在 JAX 中閉包常量是指在對一個函數(shù)進(jìn)行追蹤tracing時遇到的、不依賴該函數(shù)任何參數(shù)的非標(biāo)量數(shù)組。JAX 的jax.numpy和lax等操作是stage out的即被記錄進(jìn)計(jì)算圖而不是立即執(zhí)行因此它們不會產(chǎn)生閉包常量而原生的 NumPy 操作或預(yù)先構(gòu)造好的jax.Array則會。文檔給出了一個非常直觀的例子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預(yù)先構(gòu)造的jax.Array和np.full((16,), 42.)NumPy 原生的ndarray都是閉包常量而jnp.full((16,), 142.)是 JAX 操作在追蹤時被記錄為計(jì)算圖節(jié)點(diǎn)不是閉包常量。閉包常量為何值得警惕閉包常量最容易在不知不覺中被引入。典型場景包括在jitted函數(shù)體外預(yù)先計(jì)算好的權(quán)重矩陣、掩碼mask或索引表被函數(shù)體直接引用在函數(shù)體內(nèi)部直接調(diào)用 NumPy 函數(shù)如np.ones、np.arange這些結(jié)果會在追蹤時被物化為常量嵌入計(jì)算圖從數(shù)據(jù)加載流程中讀入的、形狀與函數(shù)參數(shù)無關(guān)的輔助數(shù)據(jù)。當(dāng)這些常量較大時它們會被內(nèi)聯(lián)進(jìn) HLO 代碼導(dǎo)致后續(xù)一系列問題詳見 Lowering 階段的取舍。使用 JAX_CAPTURED_CONSTANTS_WARN_BYTES 診斷意外捕獲文檔指出可以設(shè)置環(huán)境變量JAX_CAPTURED_CONSTANTS_WARN_BYTES為任意非負(fù)值從而在函數(shù) lowering 期間記錄警告所有不小于該字節(jié)數(shù)的閉包常量幫助你發(fā)現(xiàn)意外捕獲。從 jax/_src/config.py 的源碼可以看到該配置的真實(shí)定義captured_constants_warn_bytes int_state( namejax_captured_constants_warn_bytes, default2 * 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. ) )關(guān)鍵信息配置項(xiàng)默認(rèn)值說明jax_captured_constants_warn_bytes2 * 10 ** 9約 2GB捕獲常量總字節(jié)數(shù)超過該閾值時發(fā)出警告設(shè)為-1可徹底禁用警告jax_captured_constants_report_frames0報(bào)告中為每個捕獲常量顯示的調(diào)用棧幀數(shù)-1打印完整幀0禁用報(bào)告。注意僅當(dāng)捕獲常量總字節(jié)數(shù)超過警告閾值時才生成報(bào)告生成報(bào)告開銷較大在 jax/_src/interpreters/mlir.py 中check_jaxpr_constants與log_closed_over_constant實(shí)現(xiàn)了該警告邏輯當(dāng)closed_jaxpr.consts的nbytes總和超過閾值時warnings.warn會提示大量常量在 lowering 期間被捕獲共 N 字節(jié)并建議要么確認(rèn)這是有意的要么通過JAX_CAPTURED_CONSTANTS_WARN_BYTES-1關(guān)閉警告如需定位捕獲位置可設(shè)置JAX_CAPTURED_CONSTANTS_REPORT_FRAMES-1獲取棧幀報(bào)告。新實(shí)現(xiàn)概覽JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS文檔強(qiáng)調(diào)以下描述的是未來文檔寫作時點(diǎn)為 2026 年 4 月的常量內(nèi)部實(shí)現(xiàn)細(xì)節(jié)。它還不是當(dāng)前默認(rèn)實(shí)現(xiàn)需要通過環(huán)境變量顯式開啟JAX_USE_SIMPLIFIED_JAXPR_CONSTANTSTrue源碼 jax/_src/config.py 中對這個開關(guān)的定義佐證了這一點(diǎn)use_simplified_jaxpr_constants bool_state( namejax_use_simplified_jaxpr_constants, defaultFalse, 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_keyTrue, include_in_trace_contextTrue)注意兩點(diǎn)該 flag 的include_in_jit_keyTrue、include_in_trace_contextTrue意味著它會參與 jit 緩存鍵與追蹤上下文的構(gòu)成——不同取值下編譯出的可執(zhí)行文件不能混用緩存源碼注釋明確警告DO NOT RELY ON THIS FLAG這是一個過渡期標(biāo)志不應(yīng)在用戶代碼中長期依賴。舊實(shí)現(xiàn)的細(xì)節(jié)及其缺陷見 Previous implementation舊實(shí)現(xiàn)。Tracing 階段core.Literal 與 is_literalable當(dāng) JAX 追蹤遇到一個常量——無論它是某個 JAX primitive算子的參數(shù)還是函數(shù)的返回值——它會被表示為core.Literal并隨使用它的 primitive 一起內(nèi)嵌在Jaxpr中。決定哪些常量會被轉(zhuǎn)換為core.Literal的函數(shù)是core.is_literalable。根據(jù) jax/_src/core.py 的實(shí)現(xiàn)所有標(biāo)量常量都會被轉(zhuǎn)換為core.Literalliteralable_scalar_types走快速路徑直接返回True非標(biāo)量的np.ndarray與jax.Array也會被轉(zhuǎn)換為core.Literal當(dāng)use_simplified_jaxpr_constants開啟時jax.ArrayArrayImpl在非for_ad場景下也會字面化do_lit_array not for_ad這是為了在自動微分AD下保留常量其余類型例如自定義 Python 對象則落入選集literalable_types僅在滿足條件時字面化否則以constvars閉包變量形式出現(xiàn)在Jaxpr上。同時core.is_hoistablejax/_src/core.py判斷一個Literal是否需要被提升為參數(shù)def is_hoistable(v: Literal) - bool: return (np.ndim(v.val) 0 and getattr(v.val, nbytes, 4) config.embedded_constants_max_bytes.value)即非標(biāo)量且字節(jié)數(shù)超過embedded_constants_max_bytes的常量才值得提升小常量會被直接內(nèi)嵌見下文。Lowering 階段常量提升Hoisting為 const_args為什么不直接內(nèi)聯(lián) stablehlo.constant理論上lowering 到 HLO 時最簡單的方式是為每個core.Literal直接發(fā)射一個stablehlo.constant操作。但文檔明確列出了這樣做的一系列弊端主機(jī)內(nèi)存壓力與分片丟失如果常量是jax.Array如例子中的a_jax_arraylowering 期間會把它從設(shè)備拉回主機(jī)可執(zhí)行模塊執(zhí)行時再重新物化到設(shè)備上。這會顯著增加主機(jī)內(nèi)存占用有時是數(shù)量級的增長更進(jìn)一步如果常量在多個設(shè)備上做了分片sharding這種分片信息在拉回-重新物化的過程中會丟失。HLO 膨脹與編譯變慢大常量尤其被多次復(fù)用的同一個常量會顯著增大 HLO 體積XLA 編譯器還會嘗試對它們做常量折疊constant-folding引發(fā)告警并拖慢編譯。數(shù)值差異風(fēng)險(xiǎn)實(shí)測中 XLA 的常量折疊有時會產(chǎn)生與編譯后代碼略有不同的數(shù)值結(jié)果。jaxpr_const_args掃描并去重常量文檔指出lowering 期間使用core.jaxpr_const_args來掃描一個Jaxpr返回其中包含的常量列表按id去重uniquified。該函數(shù)對每個Jaxpr及其子Jaxpr調(diào)用結(jié)果會被記憶化memoized???jax/_src/core.py 的真實(shí)實(shí)現(xiàn)partial(weakref_lru_cache, trace_context_in_keyFalse) 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())實(shí)現(xiàn)要點(diǎn)通過weakref_lru_cache記憶化同時以id哈希為基礎(chǔ)因此同一常量不會重復(fù)掃描只收集is_hoistable非標(biāo)量、字節(jié)數(shù)超過embedded_constants_max_bytes的Literal遍歷outvars與每個方程的invars同時通過eqn_params_const_args遞歸收集方程參數(shù)中嵌套Jaxpr子函數(shù)的常量在use_simplified_jaxpr_constantsFalse默認(rèn)時直接返回空列表即舊行為不受影響。const_args 的參數(shù)排布與 const_lowering 映射所有被降級的 HLO 函數(shù)都會為Jaxpr中出現(xiàn)的每個唯一常量多接收一個額外參數(shù)。這些參數(shù)稱為const_args其排布位置是維度變量參數(shù)dimension variable args之后 → token 參數(shù)之后 → 實(shí)際數(shù)組參數(shù)array arguments之前l(fā)owering 期間維護(hù)一個映射const_lowering: dict[int, mlir.IrValues]該映射以常量的id為鍵值為對應(yīng)的 HLO 值被存放在mlir.LoweringRuleContext中。mlir.ir_constant在遇到常量時會優(yōu)先復(fù)用const_lowering中已有的 lowering而不是重新發(fā)射stablehlo.constant見 jax/_src/interpreters/mlir.py其中_ir_constant會在const_lowering命中時直接復(fù)用既有值。小常量例外embedded_constants_max_bytes存在一個例外尺寸不超過config.embedded_constants_max_bytes的小常量不會被提升為參數(shù)而是直接內(nèi)嵌embed進(jìn)生成的 HLO 與可執(zhí)行文件中。該配置定義于 jax/_src/config.pyembedded_constants_max_bytes int_state( namejax_embedded_constants_max_bytes, default32, 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_keyTrue, include_in_trace_contextTrue)默認(rèn)值為32 字節(jié)。也就是說小于等于 32 字節(jié)的非標(biāo)量常量以及所有標(biāo)量常量仍以內(nèi)聯(lián)stablehlo.constant形式存在方便 XLA 做常量折疊大于 32 字節(jié)的常量才被提升為const_args。與use_simplified_jaxpr_constants一樣它同樣參與 jit 緩存鍵與追蹤上下文。內(nèi)部函數(shù)inner function的 lowering當(dāng) lowering 一個 HLO 內(nèi)部函數(shù)非main函數(shù)時會再次調(diào)用core.jaxpr_const_args獲取對應(yīng)Jaxpr中實(shí)際的常量。這些常量預(yù)期已經(jīng)包含在外層函數(shù)的const_lowering中內(nèi)部函數(shù)會獲得自己更小的一組const_args和自己的const_lowering映射用于 lowering 其函數(shù)體。文檔舉例mlir.lower_jaxpr_as_fun就是發(fā)生此類邏輯的一處。而mlir.jaxpr_subcompjax/_src/interpreters/mlir.py不會創(chuàng)建新的 HLO 函數(shù)而是在當(dāng)前函數(shù)內(nèi)創(chuàng)建一個 block并復(fù)用外層函數(shù)的const_lowering。仍會出現(xiàn)的 stablehlo.constant文檔特別說明即便在新實(shí)現(xiàn)下降級代碼中依然會存在stablehlo.constant出現(xiàn)在以下四種場景標(biāo)量常量希望將這些常量暴露給 XLA 做常量折疊小常量尺寸不超過embedded_constants_max_bytes默認(rèn) 32 字節(jié)的常量如上文所述直接內(nèi)嵌lowering 期間新產(chǎn)生的常量常量未出現(xiàn)在被追蹤的程序中因此不在Jaxpr里。例如某些 PRNG隨機(jī)數(shù)函數(shù)的 lowering 就自帶了常量導(dǎo)出export場景目前導(dǎo)出時不提升常量參數(shù)因?yàn)閷?dǎo)出序列化尚不支持?jǐn)?shù)組序列化。這是通過mlir.LoweringParameters.hoist_constants_as_args參數(shù)控制的其默認(rèn)值與use_simplified_jaxpr_constants一致見 jax/_src/interpreters/mlir.py。avals、shardings 與 layouts 的高層計(jì)算還有一個實(shí)現(xiàn)細(xì)節(jié)部分內(nèi)部 lowering 函數(shù)需要用到參數(shù) avals有時還需要參數(shù)的 shardings 與 layouts。而且包括const_args在內(nèi)的所有參數(shù)的 avals、shardings、layouts 在 lowering 之后也仍然會被使用。因此比較方便的做法是在調(diào)用棧的較上層一次性算好例如在pxla.lower_sharding_computations中計(jì)算并向下傳遞。具體來說mlir.lower_jaxpr_to_module、pjit._pjit_cached_lower_jaxpr_to_fun、mlir.lower_jaxpr_to_fun這些函數(shù)都接收in_avals、in_shardings、in_layouts這些列表同時包含const_args的 avals 與常規(guī)參數(shù)的 avals后者對應(yīng)Jaxpr.invars此外還接收一個num_const_args參數(shù)用于區(qū)分常量參數(shù)與常規(guī)參數(shù)。編譯與執(zhí)行const_args 如何傳入可執(zhí)行文件lowering 出的 MLIR 模塊包含 const_args 對應(yīng)的參數(shù)因此編譯后的可執(zhí)行文件在被調(diào)用時也必須傳入 const_args。這里的關(guān)鍵設(shè)計(jì)問題是在哪個位置把 const_args 拼接到調(diào)用參數(shù)前面。文檔給出了一個示例強(qiáng)調(diào)第二次調(diào)用應(yīng)命中 C jit 緩存而不執(zhí)行任何 Python 代碼const jnp.array([42.]) f jax.jit(lambda: const) f() f()這意味著const必須以某種方式在 C 側(cè)傳給可執(zhí)行文件因此被存儲在pxla.MeshExecutableFastpathData中。相應(yīng)地C 緩存未命中函數(shù)例如pjit._cpp_pjit.cache_miss或pxla.MeshExecutable.create_cpp_call中的aot_cache_miss不接收 const_args 作為參數(shù)而是由這些緩存未命中函數(shù)負(fù)責(zé)自行前置拼接prependconst_args。關(guān)于 C 快速路徑fast path的支持情況從jaxlib 0.7.1開始C 快速路徑支持 const_args在更早的版本中只要存在 const_args快速路徑就會被禁用回退到較慢的 Python 路徑。const_args 在 stage 對象中的存放為實(shí)現(xiàn)上述方案const_args被保存在以下對象中stages.Loweringstages.Loweredstages.CompiledCallParamspxla.MeshExecutable注意在stages.Compiled中in_avals等字段不包含const_args即Compiled對外呈現(xiàn)的接口不含常量參數(shù)。序列化編譯緩存與 const_args一個有趣的推論是當(dāng)序列化可執(zhí)行文件例如用于編譯緩存時無需序列化閉包常量本身——可執(zhí)行文件本身不包含這些常量它只是需要接收它們作為 const_args。因此反序列化緩存的可執(zhí)行文件的一方必須自行提供 const_args。這要求編譯緩存的消費(fèi)者在緩存命中時仍能拿到與編譯時一致的閉包常量。AOT 模式與 x64 的一致性要求在 AOT預(yù)先編譯模式下lowering 與執(zhí)行可能使用不同的jax_enable_x64配置值。文檔給出約束如果常量是 64 位ndarray那么 lowering 與執(zhí)行必須使用相同的jax_enable_x64值否則常量解釋會不一致可能導(dǎo)致錯誤結(jié)果或崩潰。Previous implementation舊實(shí)現(xiàn)與缺陷當(dāng)JAX_USE_SIMPLIFIED_JAXPR_CONSTANTSFalse時即文檔寫作時點(diǎn)的默認(rèn)行為采用的是 2025 年 7 月的舊方案當(dāng) JAX 將函數(shù)追蹤成Jaxpr時會把閉包值收集進(jìn)一個常量集合并給Jaxpr加上一組對應(yīng)的constvars真正的函數(shù)參數(shù)由invars表示。大多數(shù)追蹤函數(shù)如trace_to_jaxpr_dynamic會同時返回Jaxpr和這些常量。代碼中大量使用core.ClosedJaxpr類它封裝了一個Jaxpr以及與其constvars對應(yīng)的consts。文檔明確列出了ClosedJaxpr方案的若干問題內(nèi)聯(lián)問題ClosedJaxpr中consts的 lowering 會直接產(chǎn)生內(nèi)聯(lián)的stablehlo.constant即前文描述的各種弊端主機(jī)內(nèi)存、HLO 膨脹、常量折疊數(shù)值差異、分片丟失。類型混淆Jaxpr與ClosedJaxpr在 JAX 中無處不在且常被籠統(tǒng)地命名為jaxpr難以區(qū)分當(dāng)前拿到的是哪一種。雖然已開始添加類型聲明但部分代碼仍用isinstance條件分支同時兼容兩者。緩存鍵與記憶化困難Jaxpr和ClosedJaxpr有時被用作緩存鍵且按id哈希因此希望記憶化它們的構(gòu)造。例如pe.closed_jaxpr位于 jax/_src/interpreters/partial_eval.py記憶化了ClosedJaxpr的構(gòu)造但僅在consts為空時——因?yàn)橛袝r常量不可哈希。lowering 覆蓋不全處理ClosedJaxpr中的常量需要額外小心。例如 Mosaic lowering 中尚有未實(shí)現(xiàn)非空常量ClosedJaxpr處理的地方見 jax/_src/pallas/mosaic/lowering.py 附近的相關(guān)邏輯。變換中的額外輸入將閉包常量轉(zhuǎn)成輸入后在各變換transformations中需要小心處理這些輔助輸入auxiliary inputs的傳遞。這些缺陷正是新實(shí)現(xiàn)簡化 Jaxpr 常量要解決的問題把常量顯式表示為core.Literal、統(tǒng)一通過jaxpr_const_args去重掃描并按需提升為const_args從而避免內(nèi)聯(lián)stablehlo.constant的各種問題。實(shí)踐建議與總結(jié)綜合文檔與源碼針對閉包常量可以給出如下實(shí)踐要點(diǎn)診斷先行在開發(fā)階段設(shè)置JAX_CAPTURED_CONSTANTS_WARN_BYTES如JAX_CAPTURED_CONSTANTS_WARN_BYTES1048576表示 1MB觀察是否有非預(yù)期的大常量被捕獲配合JAX_CAPTURED_CONSTANTS_REPORT_FRAMES-1獲取捕獲位置的調(diào)用棧報(bào)告。不需要時用-1關(guān)閉警告避免每次 lowering 都產(chǎn)生告警。理解參數(shù)排布在新實(shí)現(xiàn)下const_args位于維度變量參數(shù)與 token 參數(shù)之后、數(shù)組參數(shù)之前所有參數(shù)含 const_args的 avals/shardings/layouts 由調(diào)用棧上層統(tǒng)一計(jì)算并向下傳遞。緩存與序列化的約定C jit 緩存命中要求常量以const_args形式在 C 側(cè)傳遞jaxlib ≥ 0.7.1編譯緩存反序列化時不包含常量本身緩存消費(fèi)者必須自己提供 const_args。x64 一致性AOT 場景下若常量是 64 位ndarray必須保證 lowering 與執(zhí)行使用相同的jax_enable_x64。過渡期標(biāo)志JAX_USE_SIMPLIFIED_JAXPR_CONSTANTS與jax_embedded_constants_max_bytes默認(rèn) 32 字節(jié)都是過渡性配置且參與 jit 緩存鍵與追蹤上下文不應(yīng)在用戶代碼中長期依賴應(yīng)關(guān)注 JAX 版本演進(jìn)以遷移到默認(rèn)行為。本文的所有關(guān)鍵結(jié)論均可在倉庫源碼中得到印證核心邏輯 core.py、配置定義 config.py、lowering 實(shí)現(xiàn) mlir.py 與 partial_eval.py。建議讀者在閱讀本文后結(jié)合上述源碼文件與 constants.md 原文進(jìn)一步追蹤jax.jit從追蹤到執(zhí)行的完整常量處理鏈路?!久赓M(fèi)下載鏈接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more項(xiàng)目地址: https://gitcode.com/GitHub_Trending/ja/jax創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考