)
Just-in-time compilationJAX JIT 編譯原理與實戰指南【免費下載鏈接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more項目地址: https://gitcode.com/GitHub_Trending/ja/jax導讀jax.jit是 JAX 中最核心的變換之一它把一段 Python 函數在追蹤tracing階段規約成中間表示 jaxpr再交給 XLA 編譯器做即時Just-In-Time編譯最終生成針對 CPU/GPU/TPU 高度優化的可執行代碼。本篇文章圍繞 docs/jit-compilation.md 展開先講透 JAX 變換與 jaxpr 的原理再手把手演示 SELU 算子從逐 op 執行到 JIT 加速的完整過程并深入剖析「為什么不能無腦 JIT 一切」、靜態參數static_argnums/static_argnames的使用時機以及 JIT 緩存與重編譯的行為邊界。讀完你將掌握 JAX JIT 的正確打開方式并能用jax.make_jaxpr觀察函數內部到底發生了什么。JAX 變換是如何工作的JAX 允許我們變換 Python 函數其秘密在于JAX 會把每個函數規約成一系列 primitive原語操作每個 primitive 代表一個最基本的計算單元。而jax.make_jaxpr就是觀察這一過程的窗口——它返回函數的 jaxprJAX 的中間表示讓我們直觀看到追蹤結果。import jax import jax.numpy as jnp global_list [] def log2(x): global_list.append(x) ln_x jnp.log(x) ln_2 jnp.log(2.0) return ln_x / ln_2 print(jax.make_jaxpr(log2)(3.0))輸出會是一段類似下面這樣的 jaxpr{ lambda ; a:f32[]. let b:f32[] log a c:f32[] log 1.0 d:f32[] div b c in (d,) }關于 jaxpr 各字段lambda、類型、let 綁定、in 返回值的語義可以參考 docs/601/jaxpr.md 的詳細說明。關鍵點副作用不會進入 jaxpr注意上面log2函數體里有global_list.append(x)但 jaxpr 里完全沒有與它對應的內容。這并非 bug而是特性JAX 變換被設計為只理解**無副作用功能純**的代碼。如果對「純函數」「副作用」這些術語不熟悉可以在倉庫的 docs/notebooks/Common_Gotchas_in_JAX.md 的 Pure Functions 一節找到通俗解釋。不純函數在 JAX 變換下是危險的它們可能靜默失敗或者產生像Tracer 泄漏tracer leak那樣令人困惑的下游錯誤而且 JAX 往往無法自動檢測到副作用的存在。文檔給出的官方替代方案是想要調試打印用jax.debug.print其實現位于 jax/_src/debug.py想表達通用副作用、接受性能損失用jax.experimental.io_callback實現在 jax/_src/callback.py 的io_callback想檢查 tracer 泄漏、接受性能損失用jax.check_tracer_leaks。追蹤tracing的微觀機制追蹤時JAX 會用tracer 對象包裝每個參數。這些 tracer 會記錄函數調用期間發生在普通 Python 層面對它們執行的所有 JAX 操作然后 JAX 用這些記錄重建整個函數重建的產物就是 jaxpr。由于 tracer 不記錄 Python 副作用副作用自然不會出現在 jaxpr 中——但注意副作用在追蹤過程中仍然真實發生了。一個典型例子是 Python 的print()def log2_with_print(x): print(printed x:, x) ln_x jnp.log(x) ln_2 jnp.log(2.0) return ln_x / ln_2 print(jax.make_jaxpr(log2_with_print)(3.))你會發現打印出來的x是一個Traced對象——這正是 JAX 內部機制在起作用。「Python 代碼至少會被執行一次」嚴格來說是實現細節不應依賴它但在調試時可以利用這一點打印中間值。jaxpr 只反映「參數給定」的那一次執行jaxpr 捕獲的是函數在給定參數上的那次執行路徑。如果函數里有 Python 條件分支jaxpr 只會包含實際走到的那個分支def log2_if_rank_2(x): if x.ndim 2: ln_x jnp.log(x) ln_2 jnp.log(2.0) return ln_x / ln_2 else: return x print(jax.make_jaxpr(log2_if_rank_2)(jax.numpy.array([1, 2, 3])))傳入的是 1 維數組x.ndim 2為假所以 jaxpr 只會包含else分支的return x。這暗示了一個重要的推論jaxpr 的形狀/類型是由追蹤時的輸入決定的一旦輸入 shape 或 dtype 改變JAX 就不得不重新追蹤和重新編譯。從源碼看jax.make_jaxpr的實現位于 jax/_src/api.py 的make_jaxpr函數約 L2136 起。它本質上是對jit(fun, static_argnumsstatic_argnums).trace(...)的一次封裝先做一次純追蹤再把常量consts重新合并回 jaxpr 返回給調用者。其 docstring 也明確說明jaxpr 是「基于帶 let 綁定的簡單類型化一階 lambda 演算」的追蹤中間表示make_jaxpr返回的是抽象到ShapedArray級別的追蹤結果。JIT 編譯一個函數JAX 允許同一份代碼跑在 CPU/GPU/TPU 上但默認逐 op 把操作發給加速器這限制了 XLA 編譯器的優化空間。下面以深度學習常用的SELUScaled Exponential Linear Unit算子為例import jax import jax.numpy as jnp def selu(x, alpha1.67, lambda_1.05): return lambda_ * jnp.where(x 0, x, alpha * jnp.exp(x) - alpha) x jnp.arange(1000000) %timeit selu(x).block_until_ready()這段代碼的問題在于每次只把單個 op 發給加速器XLA 無法看到全局做融合優化。我們的目標自然是把盡可能多的代碼交給 XLA 編譯器讓它做整體優化。jax.jit就是為此設計的selu_jit jax.jit(selu) # 先預編譯一次再計時 selu_jit(x).block_until_ready() %timeit selu_jit(x).block_until_ready()剛才發生了什么selu_jit jax.jit(selu)得到selu的編譯版本一個被包裝的函數。調用一次selu_jit(x)JAX 在這里做追蹤——它必須有真實輸入才能用 tracer 包裝。得到的 jaxpr 再由 XLA 編譯成針對 GPU/TPU 優化過的代碼隨后立即執行以滿足這次調用。之后的每次調用都直接復用已編譯代碼完全跳過 Python 實現。如果不單獨做 warm-up基準測試會把編譯時間也計進去——雖然因為循環很多次整體仍會更快但那就不是公平對比了。計時對編譯版本測速。注意這里用了block_until_ready()這是因為 JAX 采用異步派發async dispatch需要顯式等待結果就緒再計時。關于異步派發機制可以參閱 docs/async_dispatch.rstblock_until_ready的定義見 jax/_src/api.py 與 jax/_src/array.py。值得說明的是jax.jit在 jax/_src/api.py 中的簽名還支持更多參數完整的默認值如下節選in_shardings/out_shardings輸入輸出的分片sharding規格配合jax.sharding使用static_argnums/static_argnames標記靜態編譯期常量參數donate_argnums/donate_argnames標記可捐贈的緩沖區幫助 XLA 復用輸入內存、降低峰值內存keep_unused默認False未被函數使用的參數會從編譯產物中剔除、不傳上設備device/backend顯式指定運行設備或后端cpu/gpu/tpuinline嵌套 jitted 函數的內聯策略默認False即jax.Inline.AUTOcompiler_options傳遞給 XLA 編譯器的選項字典。其中donate_argnums是文檔未展開但實踐中很重要的能力捐贈后不能再復用這些緩沖區JAX 會在你嘗試復用時報錯。更多細節可以參考倉庫中的 docs/buffer_donation.md。為什么不能無腦 JIT 一切看完上面的加速效果你可能會想干脆給所有函數都套上jax.jit得了。先來看兩個 JIT 會失敗的例子# 對 x 的「值」做條件判斷 def f(x): if x 0: return x else: return 2 * x jax.jit(f)(10) # 報錯 # 循環條件依賴 x 和 n 的值 def g(x, n): i 0 while i n: i 1 return x i jax.jit(g)(10, 20) # 報錯根因用運行時值控制追蹤期流程兩個例子的共同問題是試圖用運行時runtime值來控制追蹤期trace-time的程序流程。在 JIT 內部被追蹤的值如這里的x、n只能通過它們的靜態屬性——例如shape或dtype——來影響控制流而不能通過它們的具體數值。if x 0這樣的判斷發生在追蹤期此時x是 tracer對它做布爾比較會直接拋出 tracer 錯誤。關于 Python 控制流與 JAX 的交互細節請參閱 docs/control-flow.md。解法一改寫代碼或使用 lax 控制流應對這個問題一種方式是改寫代碼、避免對值做條件判斷另一種是使用 docs/201/control-flow.md 中介紹的特殊控制流原語比如jax.lax.cond其實現位于 jax/_src/lax/control_flow/conditionals.py。解法二只 JIT 編譯函數的一部分有時候改寫不現實那就考慮只 JIT 函數中計算最昂貴的部分。比如循環體是整個函數的熱點就只 JIT 循環體但要小心下一節講的緩存問題避免適得其反# 循環條件依賴 x 和 n但循環體是 JIT 的 jax.jit def loop_body(prev_i): return prev_i 1 def g_inner_jitted(x, n): i 0 while i n: i loop_body(i) return x i g_inner_jitted(10, 20)外層while i n仍是普通 Python 循環可以按值判斷內層loop_body被 JIT 編譯熱點計算獲得了加速。這是「部分 JIT」的典型模式。把參數標記為 static靜態參數如果確實需要 JIT 一個「對輸入值做條件判斷」的函數可以告訴 JAX對某個輸入使用抽象程度更低更具體的 tracer。方法是指定static_argnums按位置索引或static_argnames按參數名。代價是顯著的靜態參數的每個不同取值都會產生不同的 jaxpr 和編譯產物JAX 不得不為每個新值重新編譯。所以只有在該函數只會遇到有限的靜態取值集合時這才是個好策略。f_jit_correct jax.jit(f, static_argnums0) print(f_jit_correct(10))g_jit_correct jax.jit(g, static_argnames[n]) print(g_jit_correct(10, 20))以裝飾器形式使用時用裝飾器工廠模式jax.jit(static_argnames[n]) def g_jit_decorated(x, n): i 0 while i n: i 1 return x i print(g_jit_decorated(10, 20))源碼層面的約定與限制結合 jax/_src/api.py 中jit的 docstring靜態參數還有一些容易踩坑的約定靜態參數必須是可哈希的實現__hash__和__eq__且不可變因為它們的值會參與編譯緩存鍵compilation cache key的計算。文檔特別強調非數組類型或數組容器之外的參數必須標記為 static否則無法被正確追蹤。如果只給了static_argnums而沒給static_argnames或反之JAX 會用inspect.signature(fun)自動推斷對應的參數名/位置如果兩者都給了則只把顯式列出的參數當作靜態不再推斷。從 JAX v0.8.1 起jit支持省略函數參數的裝飾器工廠寫法即jax.jit(static_argnames[n])而非partial(jax.jit, ...)上面的示例正是官方推薦的現代寫法舊版本則常用functools.partial實現同樣效果。JAX 對fun持有弱引用作為緩存鍵的一部分因此fun必須可被弱引用weakly-referenceable。JIT 與緩存第一次 JIT 調用有編譯開銷所以理解jax.jit何時、如何緩存編譯結果是用好它的關鍵。緩存的基本規則假設f jax.jit(g)首次調用f時完成追蹤 編譯XLA 代碼被緩存后續調用f直接復用緩存代碼不再重復編譯——這就是jax.jit攤平編譯前期成本的方式。如果指定了static_argnums那么只有靜態參數取值與緩存一致時才復用任何一個靜態值變化都會觸發重編譯。如果靜態參數取值范圍很大程序花在編譯上的時間可能比逐 op 執行還多——這是常見的性能陷阱。不要在循環里對臨時函數調用 jit文檔明確警告避免在循環或其他 Python 作用域內對臨時函數調用jax.jit。原因在于緩存依賴函數的哈希當等價函數被反復重新定義哈希不同時緩存就失效了導致每次循環迭代都重新編譯from functools import partial def unjitted_loop_body(prev_i): return prev_i 1 def g_inner_jitted_partial(x, n): i 0 while i n: # 別這么做每次 partial 返回的函數哈希都不同 i jax.jit(partial(unjitted_loop_body))(i) return x i def g_inner_jitted_lambda(x, n): i 0 while i n: # 別這么做lambda 每次也返回哈希不同的函數 i jax.jit(lambda x: unjitted_loop_body(x))(i) return x i def g_inner_jitted_normal(x, n): i 0 while i n: # 這樣沒問題JAX 能找到緩存的編譯函數 i jax.jit(unjitted_loop_body)(i) return x i print(jit called in a loop with partials:) %timeit g_inner_jitted_partial(10, 20).block_until_ready() print(jit called in a loop with lambdas:) %timeit g_inner_jitted_lambda(10, 20).block_until_ready() print(jit called in a loop with caching:) %timeit g_inner_jitted_normal(10, 20).block_until_ready()結論很直觀partial和lambda每次都會產生新的函數對象、新的哈希緩存形同虛設而直接傳入同一個模塊級函數時JAX 能穩定命中緩存。緩存機制的底層佐證從源碼角度看JAX 在派發層大量使用帶緩存的裝飾器來復用編譯產物例如 jax/_src/dispatch.py 中的xla_primitive_callableL98 附近使用util.cache()緩存 primitive 的 callable并有util.test_event(xla_primitive_callable_cache_miss)這樣的測試探針標記緩存未命中事件同一文件還有多處util.weakref_lru_cache/util.cache(max_size2048, ...)用于緩存各類派發中間結果。這正是「同一函數對象反復jit能命中緩存、哈希不同的臨時函數會反復編譯」這一文檔結論在實現層面的體現。小結JIT 的正確打開方式盡量 JIT 大片代碼把盡可能多的計算交給 XLA讓它做算子融合與設備級優化用block_until_ready()配合%timeit得到不含異步派發誤差的公平基準。保持函數純追蹤期只記錄 primitive 操作副作用既不進 jaxpr 也會帶來 tracer 泄漏等隱患調試打印用jax.debug.print通用副作用用jax.experimental.io_callback。控制流按「靜態屬性」而非「值」走需要按值分支/循環時改寫代碼、改用jax.lax.cond等原語或只 JIT 熱點內層部分。靜態參數要克制static_argnums/static_argnames只適用于取值集合有限、可哈希的場景否則會陷入重編譯泥潭。把jit用在穩定、可緩存的函數對象上避免在循環內用partial/lambda現造函數再 JIT。如果需要更系統地學習可以繼續閱讀倉庫中 docs/201/jit.md 關于jit進階語義、docs/601/jaxpr.md 關于 jaxpr 語言以及 docs/async_dispatch.rst 關于異步派發的說明。【免費下載鏈接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more項目地址: https://gitcode.com/GitHub_Trending/ja/jax創作聲明:本文部分內容由AI輔助生成(AIGC),僅供參考