模型層 開放閱讀

JAX

JAX

概念 ID
jax
更新時間
2026-05-29
來源數量
待補

JAX

3 秒看懂

JAX 是一個由 Google Brain 團隊開發的、基於 Python 的高效能數值計算庫和自動微分系統。它通過函數語言程式設計範式、即時(JIT)編譯至 XLA 以及自動向量化/並行化,將 Python 的靈活性與底層硬體(如 GPU/TPU)的極致效能相結合,主要面向 AI 及科學計算領域的前沿研究。

3 分鐘產業解釋

JAX 的出現,本質上是為了解決“研究靈活性與生產級效能”之間的長期矛盾。在傳統 AI 研發中,研究人員通常使用 PyTorch 等易於除錯的動態圖架構快速迭代想法,但在需要大規模訓練或部署時,又不得不將程式碼移植到 TensorFlow 等靜態圖架構以追求效能,這個過程耗時且容易出錯。

JAX 的核心價值在於**“研究到生產”路徑的一體化**。它允許研究人員用純 Python 和 NumPy 風格的程式碼寫出複雜的模型和演算法,而 JAX 的轉換器(@jit, vmap, pmap)能在編譯階段對程式碼進行深度最佳化,並將其轉換為高效的 XLA 程式,直接在 GPU 或 TPU 上執行。這意味著,同一個程式碼庫既能用於快速實驗,也能直接擴充套件到數千個晶片上進行訓練,極大地提升了研發效率。目前,JAX 已成為 Google 內部 DeepMind 等頂尖實驗室進行前沿模型(如 Gemini、AlphaFold)開發的預設選擇之一,並在學術界和工業界(尤其是高效能運算和大規模模型領域)獲得了快速增長的影響力。

15 分鐘專家深入

要理解 JAX 的行業地位,必須將其置於與 PyTorch 和 TensorFlow 的“三元博弈”中。PyTorch 憑藉其動態圖和易用性統治了學術研究和快速原型開發;TensorFlow 憑藉其成熟的生態系統(TF Serving, TF Lite)和在生產端的廣泛部署,仍佔據工業界主流。

JAX 的技術路線差異在於其徹底的函數語言程式設計核心。所有操作都被視為無副作用的純函式,狀態(如模型引數)必須顯式傳遞。這種範式最初提高了學習門檻,但卻帶來了根本性優勢:

  1. 透明的轉換:因為函式是純的,JAX 可以安全地對其應用一系列轉換,如 JIT 編譯(融合操作、減少記憶體訪問)、自動微分(grad)和自動向量化(vmap)。vmap 可以自動將批處理維度新增到函式中,而無需手動編寫批處理程式碼。
  2. 可組合性與硬體抽象:這些轉換器可以自由組合(例如 jax.pmap(jax.jit(...)))。pmap 可以輕鬆地將一個函式在多個裝置(如 GPU)上進行單程式多資料(SPMD)並行化,這使得編寫分散式訓練程式碼比在 PyTorch 中使用 DistributedDataParallelFullyShardedDataParallel 更加簡潔。
  3. 確定性執行:函式式範式使得程式的執行更加可預測,這對於在大規模分散式系統中除錯和復現問題至關重要。

其挑戰主要在於生態。儘管發展迅速,JAX 的第三方庫、教程、社群支援和部署工具鏈仍遠不及 PyTorch 和 TensorFlow。因此,目前 JAX 更適用於對計算效能、並行化和演算法研究有極致要求的場景,而非通用的端到端應用開發。

技術原理

JAX 的技術堆疊建立在幾個關鍵元件之上,其效能優勢來源於編譯器的深度最佳化。

[使用者 Python 程式碼 (NumPy-like API)]
            |
            | (使用 @jit, vmap, pmap 等轉換器)
            v
[JAX 追蹤 (Tracing) 與 函式轉換]
            |
            | (生成計算圖/XLA HLO)
            v
[XLA 編譯器 (Accelerated Linear Algebra)]
            |
            | (針對目標硬體進行深度最佳化)
            v
[目標硬體執行 (CPU, GPU, TPU)]
  1. 函式式核心與轉換器

    • jax.jit:通過追蹤函式執行,將其轉換為 XLA(HLO)計算圖。XLA 會進行運算元融合(將多個操作合併為一個核心,減少記憶體讀寫)、記憶體版面配置最佳化等,從而顯著提升計算密集型任務的效能。
    • jax.grad:基於自動微分(反向模式)計算函式梯度。由於其函式式特性,微分可以作用於任意 Python 函式,並且可以與其他轉換器組合。
    • jax.vmap (Vectorizing Map):自動將一個處理單個樣本的函式,轉換為一個能夠處理一批樣本的函式。它通過在邏輯上新增一個批次維度並自動處理廣播和規約來實現,消除了手動編寫向量化程式碼的負擔。
    • jax.pmap (Parallelizing Map):將一個函式在多個裝置(如多個 GPU 核心)上並行執行。使用者只需編寫處理單個裝置上資料的程式碼,pmap 會自動將資料分發到各裝置並複製函式副本,但梯度聚合(如 All-Reduce)需要使用者顯式呼叫 jax.lax.pmeanjax.lax.psum 等集合通訊操作。
  2. XLA 編譯器: XLA 是效能的“秘密武器”。它不僅僅是簡單的編譯器,而是具備目的碼生成能力的領域特定編譯器。針對 GPU,它會生成高度最佳化的 CUDA 核心;針對 TPU,它會生成高效的 TPU 指令。XLA 的最佳化是跨運算元、全域性性的,這是其能實現高效能的關鍵。

  3. 硬體後端: JAX 通過 XLA 支援 CPU、NVIDIA GPU 和 Google TPU。尤其是與 TPU 的深度整合,使其成為在 TPU Pod 上進行超大規模模型訓練的首選工具,能充分發揮 TPU 的互聯和矩陣計算優勢。

技術演進史

JAX 的發展是 Google 內部對不同機器學習架構(如早期的 DistBelief、TensorFlow)經驗總結與反思的結果。

  • 起源:其前身是 Google 在 2010 年代中期開發的自動微分庫 Autograd。Autograd 的核心作者(如 Dougal Maclaurin, Matthew Johnson 等)後來加入了 Google Brain,並在此基礎上建置了 JAX。
  • 內部孕育:2018 年,JAX 專案在 Google 內部啟動,旨在建立一個結合了 NumPy 易用性、Autograd 微分能力、以及 XLA 編譯效能的庫。
  • 開源與迭代:2018 年底,JAX 在 GitHub 上開源。隨後,Google 的多個頂級研究專案(如 Trax 用於序列模型、Flax 作為其推薦的神經網路庫、T5XMaxText)都基於 JAX 建置,這極大地推動了其內部應用和生態發展。
  • 生態擴張:近年來,JAX 的影響力溢位到學術界和工業界。知名的科學計算專案(如 JAX-MD 用於分子動力學)和開源模型庫(如 DeepMind 開發的 Optax 最佳化器庫、Google/DeepMind 開發的 Orbax 檢查點庫)都採用 JAX。DeepMind 的眾多突破性工作,如 AlphaFold 2 的訓練部分、AlphaCode 等,均使用 JAX。

技術路線對比

特性JAXPyTorchTensorFlow (2.x)
程式設計範式函式式優先 (純函式,顯式狀態管理)物件導向/命令式 (動態圖,易於除錯)混合式 (Eager執行 + tf.function靜態圖)
核心優勢可組合的轉換 (JIT, vmap, pmap),編譯器深度最佳化,並行化簡潔靈活性,易用性,龐大的社群與生態,動態圖除錯方便生產部署成熟 (TF Serving, TFLite),跨平台支援廣,Keras API友好
效能來源XLA 編譯器全域性最佳化,與TPU深度整合TorchScript/AOT編譯,與NVIDIA CUDA/cuDNN庫深度整合XLA/TFRT (新執行時),計算圖最佳化
並行化方式jax.pmap (簡潔的SPMD), jax.sharding (用於大規模模型分片)torch.nn.DataParallel/DistributedDataParalleltorch.distributed.fsdptf.distribute.Strategy (MirroredStrategy, TPUStrategy等)
除錯體驗較難(需適應函式式,使用 jax.debug 工具)最佳 (標準Python偵錯程式即用)良好 (Eager模式下易除錯)
主要應用領域大規模前沿研究 (LLM, 科學計算), Google內部核心研發學術研究, 快速原型開發, 大量工業應用工業生產部署, 行動端/邊緣端, 推薦系統等
生態成熟度快速增長中,但相對較新最成熟, 庫、教程、社群支援最豐富成熟, 但新專案向PyTorch/JAX轉移趨勢明顯

上下游

上游:

  • 程式語言與庫:Python, NumPy(JAX 旨在成為 NumPy 的可微分、可加速替代品)。
  • 編譯器:XLA 是 JAX 效能的基石。
  • 硬體:NVIDIA GPU (A100, H100), Google TPU (v4, v5e), CPU。

下游:

  • 架構層
    • 神經網路庫Flax (Google 主推的 JAX NN 庫), Haiku (DeepMind 開發), Trax
    • 最佳化器庫Optax
    • 檢查點庫Orbax
    • 強化學習庫dm-haikuRLax
  • 應用層
    • 大規模語言模型:Google 的 PaLM、Gemini 系列模型的訓練架構。
    • 科學計算與物理模擬:分子動力學、流體模擬、氣象預測。
    • 計算機視覺:ViT 等模型的實現與研究。
    • 生物資訊學:蛋白質結構預測 (AlphaFold)。
  • 基礎設施/雲端服務Google Cloud TPU 是 JAX 最理想和深度最佳化的執行平台。各大雲端廠商(AWS, Azure)的 GPU 例項也支援 JAX。

關鍵指標

  • 效能:在相同硬體上,對於大型、可融合的模型,JAX + XLA 常能展現出優於 PyTorch 原生 eager 模式的吞吐量,尤其在 TPU 上優勢明顯。具體提升幅度取決於模型和硬體,需基準測試。
  • 易用性/開發效率:對熟悉函數語言程式設計和 NumPy 的研究人員友好。vmap/pmap 極大地簡化了並行化程式碼。學習曲線比 PyTorch 陡峭。
  • 可擴充套件性pmap 和新的 jax.sharding API 使得在數千個裝置上進行資料並行和模型並行訓練變得相對直接,是其核心優勢之一。
  • 編譯時間:首次執行或追蹤新函式時,JIT 編譯可能需要較長時間。後續呼叫則直接執行編譯後的程式碼。
  • 生態規模:第三方庫數量、社群問題解答速度、現成模型數量等,仍在追趕 PyTorch。

供需與市場資料

JAX 本身是開源免費軟體,其“市場”體現在使用者規模和影響力上。

  • 需求端:需求主要來自進行前沿演算法研究大規模模型訓練以及高效能科學計算的團隊。這些使用者對效能、並行化效率和可復現性的要求高於對生態豐富度的需求。Google、DeepMind 以及眾多頂尖高校和 AI 實驗室是其核心使用者。
  • 供給端:供給方主要是 Google(持續投入開發)、圍繞 JAX 的開源社群貢獻者,以及基於 JAX 提供雲端 TPU 服務的 Google Cloud。
  • 市場資料:缺乏統一的公開市場報告。可觀察的指標包括:GitHub 星標數、Stack Overflow 上相關問題的增長、學術論文中致謝使用 JAX 的比例、以及雲端廠商 TPU 服務的使用量增長(部分作為代理指標)。總體來看,JAX 的採用率呈現快速上升趨勢,尤其在高效能運算和大型模型領域,但整體市場份額仍遠小於 PyTorch 和 TensorFlow。

代表公司與資本對映

  • 核心推動者
    • Google (Alphabet):JAX 的創造者和最大使用者。JAX 是 Google Cloud TPU 生態的關鍵一環,為雲端服務帶來高價值客戶。DeepMind 的科研成果(JAX是其基石)也提升了 Alphabet 的科技形象。
    • NVIDIA:JAX 在 GPU 上執行良好,促進了 GPU 在大規模 AI 訓練中的銷售。JAX 對 CUDA 和 cuDNN 的依賴鞏固了 NVIDIA 的軟體生態護城河。
  • 重要採用者與生態夥伴
    • Hugging Face:在其 Transformers 等庫中集成了 JAX 支援,而 Optax(DeepMind 開發)、Orbax(Google/DeepMind 開發)等庫同樣原生支援 JAX,共同連線了 JAX 生態與龐大的開源模型社群。
    • 眾多AI初創公司:尤其是專注於AI for Science(如材料發現、藥物研發)、大規模生成式模型高效能運算的初創公司,將 JAX 作為核心研發架構。
  • 資本對映邏輯:投資於深度依賴 JAX 進行前沿研究的公司(如某些 AI 製藥公司),或關注 Google Cloud TPU 業務增長的投資者,需要理解 JAX 的技術優勢。JAX 的流行是高效能 AI 基礎設施需求增長的一個訊號,利好算力(晶片、雲端)提供商。

投資邏輯

  1. “研究-生產”一體化趨勢:JAX 代表的技術路線,如果最終被廣泛接受,將改變 AI 研發範式,縮短創新週期。投資應關注那些能夠最快將前沿研究轉化為產品(特別是依賴大規模模型)的公司,而 JAX 可能是其技術棧的一部分。
  2. Google AI 生態的護城河:JAX 與 TPU、Google Cloud 的深度繫結,增強了 Google 在 AI 基礎設施領域的競爭壁壘。投資 Google 或 Google Cloud 的增長,部分邏輯在於其一體化軟硬體(JAX-TPU-GCP)帶來的效能優勢和客戶粘性。
  3. AI for Science 的槓桿:JAX 在科學計算領域的統治力日益增強。投資於 AI for Science 賽道時,被投企業是否採用 JAX 可以作為一個側面的技術判斷指標,表明其追求計算效率和前沿性的決心。
  4. 風險:JAX 生態的成熟度是其主要風險。如果 PyTorch 的並行化和編譯工具(如 torch.compile)持續快速進步,可能會侵蝕 JAX 的獨特優勢。投資需觀察社群和工業界的實際採納趨勢。

常見誤讀糾偏

  1. 誤讀:JAX 就是 Google 的 PyTorch,只能用於神經網路。

    • 糾偏:JAX 是一個通用的數值計算和自動微分庫,其核心是 NumPy-like 的 API 和函式式轉換。神經網路只是其應用之一。它廣泛用於物理模擬、最佳化問題、機率程式設計等科學計算領域。將其侷限於神經網路架構是嚴重的低估。
  2. 誤讀:JAX 純函數語言程式設計意味著不能有任何狀態(如模型引數、最佳化器狀態)。

    • 糾偏:JAX 要求狀態是顯式管理的,而非禁止。常見的模式是將所有狀態(引數、最佳化器狀態)打包成一個 Python 物件(如字典、元組或自定義資料類),並作為函式的輸入和輸出傳遞。這雖然增加了樣板程式碼,但使得狀態變化清晰可追蹤,與 jitgrad 等轉換完美相容。Flax 等庫提供了更高階的模組系統來簡化這一過程。
  3. 誤讀:JAX 除錯非常困難,不實用。

    • 糾偏:除錯 jit 編譯後的程式碼確實需要特殊工具(如 jax.debug.print),因為標準偵錯程式無法介入編譯後的圖執行。然而,JAX 允許在非 jit 模式下執行程式碼進行逐行除錯。此外,jax.debug 模組提供了在編譯程式碼中注入列印和斷言的能力。除錯體驗不如 PyTorch 直觀,但通過正確的方法是可以管理的,並非不可用。

學習路徑

  1. 基礎:牢固掌握 Python 和 NumPy。理解陣列操作、廣播、索引等。
  2. 核心概念:學習 JAX 官方教程,重點理解其與 NumPy 的異同、即時編譯 (jit)自動微分 (grad)自動向量化 (vmap) 的概念與用法。
  3. 函式式思維:練習用純函式編寫程式碼,學習顯式傳遞和返回狀態。
  4. 並行化:掌握 pmap 進行多裝置資料並行,瞭解新的 jax.sharding API 用於更復雜的模型並行。
  5. 神經網路庫:選擇一個高階庫深入學習,推薦 Flax (官方推薦) 或 Haiku。學習如何用它們定義模組、管理引數。
  6. 實踐專案:在小型專案(如 MNIST)上用純 JAX + Flax 從頭實現。然後嘗試在 Colab 的 TPU 或 GPU 上使用 pmap 加速訓練。
  7. 深入:閱讀 XLA 和 HLO 相關文件,理解編譯最佳化原理。閱讀頂尖實驗室(如 DeepMind)的 JAX 專案程式碼。

一句話總結

JAX 是面向 AI 與高效能運算未來的“編譯器驅動的 Python 科學計算棧”,它以函數語言程式設計的嚴謹性,通過可組合的編譯器轉換,打通了從研究原型到大規模高效能生產程式碼的瓶頸。

延伸閱讀與來源

source: 公開揭露與公開資料整理 本頁僅用於產業鏈學習、資訊檢索和研究輔助;不構成投資建議,不預測漲跌,不提供買賣、部位或目標價建議。
完整概念頁 複盤 13 節結構 公司投研頁 沿產業鏈找到受益公司 投資課 把概念轉成可跟蹤模型