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