函式庫 / SDK
jax-ml/jax avatar
jax-ml/jax

JAX:把求导、编译與批处理组合进 NumPy 程序

Python+NumPy 程式的可組合轉換:微分、向量化、JIT 到 GPU/TPU 等。

36,304 個 Star3,779 個 ForkPythonApache-2.0

秒懂

它是什麼?
围绕 grad、jit、vmap 和设备分片,拆解 jax-ml/jax 的数值计算模型、平台边界與研究專案属性。
適合誰用?
JAX 適合已经熟悉 Python 和 NumPy、需要把同一套数值函数用于求导、编译和批处理的研究與工程团队。采用前应先按目标硬件执行 README 中对应的安裝命令,並用实际函数檢查 Python 控制流在 jit 下的行為;README 將它定义為研究專案,不能把示例直接当成生产效能或兼容性承诺。
可以商用嗎?
可以。Apache-2.0 是寬鬆授權:你可以使用、修改並販售以它為基礎的軟體,只需保留著作權與授權聲明。
還在維護嗎?
有在維護。儲存庫在最近一天內有新的提交。
用什麼語言寫的?
主要是 Python(依據 GitHub 的語言統計)。

以上回答依據專案的 GitHub 資料(最近同步於 2026年9月15日)與我們的分析,不構成法律意見。

開源專案深度解析

一個用于变换数值函数的库 · jax-ml-jax-deep-analysis

JAX 是一個面向加速器数组计算和程序变换的 Python 库,目标是為高效能数值计算和大规模机器学习服務。README 的核心表述是:JAX 可以自动对原生 Python 和 NumPy 函数求微分,並且這一能力可以穿透循环、分支、递归和闭包。反向模式微分透過 jax.grad 提供,也支持前向模式微分,两者可以以任意阶数组合。底层由 XLA 负责在 TPU、GPU 和其他硬件加速器上编译並執行 NumPy 程序。README 還指出,编译與自动微分可以任意组合。

核心的三個变换:grad、jit 與 vmap · jax-ml-jax-deep-analysis

README 將三個变换作為系统的核心来介绍。jax.grad 计算反向模式梯度,示例展示了对 tanh 函数求一阶和三阶导数,也展示了对带 if/else 分支的函数求导。jax.jit 用 XLA 对函数做端到端编译,既可以作為装饰器,也可以作為高阶函数使用,同時会限制函数内可用的 Python 控制流类型。jax.vmap 沿数组轴映射函数,將循环下推到函数的原语操作上,使矩阵向量乘法变成矩阵矩阵乘法,從而免去在代码裡显式携带批维度的麻烦。README 引导读者查阅 Autodiff Cookbook 和参考文档获取更多內容。

在同一個损失函数上组合 grad、jit 和 vmap · jax-ml-jax-deep-analysis

README 的主示例定义了一個预测函数,其中用循环遍历参数对,然后定义平方误差损失,再构造 grad_loss = jax.jit(jax.grad(loss)) 和 perex_grads = jax.jit(jax.vmap(grad_loss, in_axes=(None, 0, 0))),用于快速计算逐样本梯度。同样的模式在扩展一节再次出现:对损失做 jit 和 grad 之后,在显式分片的参数和資料上求梯度。示例输出打印出 f32[512@data,512] 這样的类型,表明参数张量带有 data 轴的分片标注。

三种扩展模式:從自动到手动 · jax-ml-jax-deep-analysis

為了把计算扩展到数千台设备,README 記錄了三种做法。编译器驱动的自动並行化讓使用者像在单台全局机器上那样编程,由编译器选择如何分片資料和划分计算,僅辅以使用者提供的一些约束。显式分片加自动划分仍保留全局视角,但資料分片在 JAX 类型中可见,可以用 jax.typeof 檢查。手动逐设备编程则切换到逐设备视角,並允许使用显式集合通信。README 中的表格按视图类型、是否显式分片、是否显式集合通信三项来概括這三种模式。

平台支持與文档中的安裝命令 · jax-ml-jax-deep-analysis

README 的平台矩阵列出了 CPU 在 Linux x86_64、Linux aarch64、Mac aarch64、Windows x86_64 和 Windows WSL2 x86_64 上的支持。NVIDIA GPU 在两种 Linux 变体上受支持,在 Windows WSL2 上為实验性,其他平台不可用。Google TPU 僅限 Linux x86_64。AMD GPU 在 Linux x86_64 上受支持,在 Windows WSL2 上為实验性。Apple GPU 在 Mac aarch64 上為实验性,Intel GPU 在 Linux x86_64 上為实验性。README 给出的安裝命令是:CPU 用 pip install -U jax,NVIDIA GPU 用 pip install -U "jax[cuda13]",Google TPU 用 pip install -U "jax[tpu]",Linux 上的 AMD GPU 用 pip install -U "jax[rocm7-local]"。Intel GPU 安裝需遵循 Intel 自己的说明。README 將源码编译、Docker、其他 CUDA 版本和社区 conda 建置等替代方案指向文档。

專案状态、引用與 Apache-2.0 许可 · jax-ml-jax-deep-analysis

README 明确说明 JAX 是一個研究專案,不是谷歌的官方产品,並提醒存在锐利边缘,指引使用者阅读 gotchas 笔记本和问题跟踪器。引用条目按字母顺序列出作者,版本号取 jax/version.py 中的 0.3.13,年份為 2018,即專案开源發布的時间。一個僅支持自动微分和 XLA 编译的早期版本曾在 SysML 2018 论文中描述;README 说目前正在撰写更全面、更新的论文。儲存庫采用 Apache-2.0 许可,授予永久、全球、非独占、免费、免版税且不可撤销的版权许可,以及相应的专利许可,该专利许可在受许可方提起专利诉讼時终止。许可摘录没有涉及保修、支持或安全态势。儲存庫元資料在本文准备時显示 36103 個星标、3716 個复刻和 2554 個开放问题。

jax-ml-jax-deep-analysis 專案核對記錄第1項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第2項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第3項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第4項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第5項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第6項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第7項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第8項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第9項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第10項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第11項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第12項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

jax-ml-jax-deep-analysis 專案核對記錄第13項:以目前倉庫中已列出的功能和檔案為準,檢查輸入是否被接受、輸出是否落在預期位置、錯誤是否能由日誌辨識,以及重跑後訂單、商品、資料庫或命令結果是否維持一致。這項記錄只針對 jax-ml-jax-deep-analysis 的實際流程,文件沒有說明的性能、相容性和安全結果不作延伸判斷。

編輯結論

JAX 適合已经熟悉 Python 和 NumPy、需要把同一套数值函数用于求导、编译和批处理的研究與工程团队。采用前应先按目标硬件执行 README 中对应的安裝命令,並用实际函数檢查 Python 控制流在 jit 下的行為;README 將它定义為研究專案,不能把示例直接当成生产效能或兼容性承诺。

官方來源

  1. Official documentation
  2. Official README
  3. Project repository
  4. Release notes
社群筆記

社群筆記