jax-ml/jax

Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more

解決的問題

JAX 提供了一個高性能的數值計算與大規模機器學習系統。開發者可以撰寫類似 NumPy 的程式碼,這些程式碼能自動微分、編譯以提升速度,並在 GPU 和 TPU 等多種硬體加速器上進行擴展。

工作原理

JAX 使用 XLA(加速線性代數)將 Python 與 NumPy 函數編譯為最佳化的機器碼。它作為一個可組合的函數轉換系統運作:

  • jax.grad:使用自動微分(支援反向模式與前向模式)計算 Python 與 NumPy 函數的梯度。
  • jax.jit:使用 XLA 將函數端對端編譯,以實現更快的執行速度。
  • jax.vmap:自動向量化函數,將其映射至陣列軸上,從而消除手動批次迴圈。

適用對象

專為從事高性能數值計算與大規模機器學習的研究人員與開發者設計,適用於需要高效梯度計算以及在數千個裝置上擴展計算能力的使用者。

主要特色

  • 可組合的轉換:可任意順序組合 gradjitvmap,創造高度優化的函數。
  • 硬體加速:透過 XLA 原生支援 NVIDIA GPU、Google TPU 及其他加速器。
  • 彈性微分能力:可對迴圈、分支、遞迴與閉包進行任意階數的微分。
  • 多種擴展方式:提供三種擴展模式:編譯器基礎的自動平行化、顯式分片(sharding)與手動設備級程式設計。

相關

  • 專案
  • 專案
  • 專案
  • 專案
  • 專案