ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

JAX 安装完全指南:CPU、NVIDIA GPU、AMD ROCm、TPU 与多平台部署详解

2026/9/10 23:21:03 拓冰建站 浏览量
JAX 安装完全指南:CPU、NVIDIA GPU、AMD ROCm、TPU 与多平台部署详解 JAX 安装完全指南CPU、NVIDIA GPU、AMD ROCm、TPU 与多平台部署详解【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxJAX 是一套对 Python NumPy 程序进行组合变换自动微分、向量化、JIT 编译到 GPU/TPU的框架。要正确使用它必须先理解其独特的双包安装架构jax纯 Python、跨平台负责 API 与变换逻辑jaxlib包含编译后的二进制与平台/加速器后端决定你能在什么硬件上运行。本文以仓库官方安装文档 docs/installation.md 为主体结合 setup.py、jax/_src/lib/init.py、jax_plugins/cuda/init.py 等源码完整讲解 CPU、NVIDIA GPUCUDA 12/13、AMD GPUROCm、Google Cloud TPU、Intel GPU实验性 oneAPI等场景的安装命令、版本兼容矩阵、底层校验机制与故障排查方法。读完本文你将能够根据自身硬件与 CUDA/ROCm 环境选择正确的安装方式并定位常见安装问题。一、理解 JAX 的双包架构jax 与 jaxlib使用 JAX 需要安装两个包jax纯 Python 实现跨平台通用包含jax.numpy、jax.grad、jax.jit、jax.vmap等全部用户层 API 与组合变换逻辑jaxlib包含编译后的二进制C 实现的 XLA 运行时、PJRT 插件、CPU 特性检测等针对不同操作系统与加速器需要不同的构建产物。这一分工在仓库 setup.py 中有直接体现install_requires中固定声明jaxlib 0.11.1, 0.11.2版本号来自 jax/version.py 中的_minimum_jaxlib_version 0.11.1与_version 0.11.2同时要求python_requires3.12并依赖numpy2.1、ml_dtypes0.5.0、opt_einsum、scipy1.15。而jax与jaxlib的版本必须相互匹配。导入 JAX 时jax/_src/lib/init.py 中的check_jaxlib_version会执行严格的版本校验若jaxlib版本低于jax要求的_minimum_jaxlib_version会抛出RuntimeError提示jaxlib is version X, but this version of jax requires version Y若jaxlib版本高于jax版本例如只升级了 jaxlib同样会报错提示二者不兼容。这也是为什么官方安装文档反复强调始终使用pip install -U jax[xxx]这种带 extra 的完整安装命令让 pip 同时解析jax、jaxlib与对应后端插件而不是手动单独安装jaxlib。二、快速安装摘要四种典型场景对于大多数用户典型的安装命令如下CPU 版Linux / macOS / Windowspip install -U jaxNVIDIA GPUCUDA 13pip install -U jax[cuda13]AMD GPUROCmpip install -U jax[rocm7-local]Google Cloud TPUTPU VMpip install -U jax[tpu]这些 extra 的具体依赖在 setup.py 的extras_require中定义例如cuda13会安装jax-cuda13-plugin[with-cuda]内含 CUDA 13 相关的 NVIDIA 官方 wheel 依赖见 jax_plugins/cuda/plugin_setup.pytpu会安装jaxliblibtpurequests后者用于jax.distributed.initialize。后面各节将逐一展开每种场景的细节。三、支持平台总览下表列出了所有受支持平台及对应的安装路径安装前请先确认自己的环境属于哪一档Linux, x86_64Linux, aarch64Mac, aarch64Windows, x86_64Windows WSL2, x86_64CPU支持支持支持支持支持NVIDIA GPU支持支持不支持不支持实验性Google Cloud TPU支持不支持不支持不支持不支持AMD GPU支持不支持不支持不支持实验性Apple GPU不支持不支持实验性不支持不支持Intel GPU实验性不支持不支持不支持不支持四、CPU 安装4.1 pip 安装CPUJAX 团队当前为以下操作系统与架构发布jaxlibCPU wheelLinux, x86_64Linux, aarch64macOSApple ARM 架构Windows, x86_64实验性在笔记本上做本地开发时CPU 版就足够了pip install --upgrade pip pip install --upgrade jaxWindows 注意如果机器上尚未安装 Microsoft Visual Studio 2019 RedistributableVC 运行库需要先安装它否则运行时可能因缺少 DLL 而失败。4.2 其他平台必须从源码构建除上述平台外其余操作系统与架构需要从源码构建。直接对不支持的系统执行pip install jax可能出现的现象是jax本体安装成功但jaxlib未能一起装上——此时导入jax会在运行时失败。这是因为 jax/_src/lib/init.py 在导入时直接import jaxlib缺失会抛出ModuleNotFoundError并提示先完成安装。从源码构建的详细流程可参考 docs/building_on_jax.md。五、NVIDIA GPU 安装CUDA5.1 硬件与驱动要求JAX 对 NVIDIA GPU 的算力SM版本有硬性要求CUDA 12支持 SM 5.2Maxwell或更新的 GPUKepler 系列已不再支持NVIDIA 自身也已停止对 Kepler 的软件支持CUDA 13支持 SM 7.5 或更新的 GPUCUDA 13 中 NVIDIA 放弃了对更老 GPU 的支持。驱动方面必须首先安装 NVIDIA 驱动。建议安装 NVIDIA 提供的最新驱动且版本下限为——Linux 上 CUDA 12 要求驱动 525CUDA 13 要求驱动 580。如果需要在较老驱动上使用更新的 CUDA 工具包例如集群环境不便升级驱动可以尝试 NVIDIA 提供的 CUDA forward compatibility packages前向兼容包。5.2 方式一推荐用 pip wheels 安装 CUDA 与 cuDNNJAX 团队强烈推荐用 pip wheel 方式安装 CUDA 与 cuDNN因为它简单得多。注意NVIDIA 只发布了 x86_64 与 aarch64 的 CUDA 包。pip install --upgrade pip # NVIDIA CUDA 13 安装 # 注意wheel 仅适用于 Linux pip install --upgrade jax[cuda13] # 或者使用 CUDA 12 # pip install --upgrade jax[cuda12]官方建议尽早迁移到 CUDA 13 wheel未来某个时间点将放弃 CUDA 12 支持。cuda即 cuda12与cuda13两个 extra 在 setup.py 中的差别仅在于安装jax-cuda12-plugin[with-cuda]还是jax-cuda13-plugin[with-cuda][with-cuda]会拉入对应的nvidia-cublas、nvidia-cudnn、nvidia-nccl、nvidia-cuda-nvcc、nvidia-cufft、nvidia-cusolver、nvidia-cusparse、nvidia-nvjitlink、nvidia-cuda-nvrtc、nvidia-nvshmem等官方库CUDA 13 还会额外包含nvidia-nvvm用于提供 libdevice见 jax_plugins/cuda/plugin_setup.py。如果 JAX 检测到错误的 CUDA 库版本需要检查两点确保没有设置LD_LIBRARY_PATH——它可能覆盖优先于pip 安装的 NVIDIA CUDA 库确保已安装的 CUDA 库确实是 JAX 要求的版本重新运行上述安装命令通常即可解决。5.3 方式二较难使用本地预装 CUDA/cuDNN如果希望使用系统里预装的 NVIDIA CUDA需要先自行安装 CUDA 与 cuDNN。JAX 只为Linux x86_64 与 Linux aarch64提供预编译 CUDA wheel其他操作系统与架构组合需要从源码构建。驱动版本应不低于 CUDA 工具包对应要求的驱动版本同样新旧驱动与新版工具包的组合问题可借助 CUDA forward compatibility packages 解决。JAX 当前提供两种 CUDA wheel 变体CUDA 12 wheel构建基于兼容范围CUDA 12.3CUDA 12.1cuDNN 9.10cuDNN 9.10.2, 10.0NCCL 2.19NCCL 2.18CUDA 13 wheel构建基于兼容范围CUDA 13.0CUDA 13.0cuDNN 9.12cuDNN 9.12, 10.0NCCL 2.19NCCL 2.18安装命令-local表示不使用 NVIDIA 的 pip wheel而是查找本地安装的 CUDA/cuDNNpip install --upgrade pip # 安装与 NVIDIA CUDA 13 及 cuDNN 9.12 或更新兼容的 wheel # 注意wheel 仅适用于 Linux pip install --upgrade jax[cuda13-local] # 安装与 NVIDIA CUDA 12 及 cuDNN 9.10 或更新兼容的 wheel # 注意wheel 仅适用于 Linux # pip install --upgrade jax[cuda12-local]这些 pip 安装方式在 Windows 上不工作且可能静默失败请对照上文支持平台总览表确认。查看自己的 CUDA 版本nvcc --versionJAX 通过LD_LIBRARY_PATH查找 CUDA 库、通过PATH查找二进制ptxas、nvlink请确保这些路径指向正确的 CUDA 安装。另外 JAX 需要libdevice10.bc通常来自cuda-nvvm包请确认 CUDA 安装中包含它。5.4 底层机制CUDA 版本自动校验与 JAX_SKIP_CUDA_CONSTRAINTS_CHECKJAX 在导入时会自动检查 CUDA 各组件的版本若库版本不够新会直接报错。这一逻辑实现在 jax_plugins/cuda/init.py 的_check_cuda_versions中它会逐一校验 CUDA runtime、cuDNN、cuFFT、cuPTI、cuBLAS、cuSPARSE 等组件的运行时版本与构建版本例如 CUDA 最低 12.10、cuBLAS 最低 12.10.0、cuPTI 最低 18 等并输出形如Version JAX was built against / Minimum supported / Installed version的诊断信息除版本检查外它还会针对已知问题给出警告例如 cuBLAS 13.2 在并发流场景下的 TMEM 释放缺陷、cuDNN 9.10.0 的二进制兼容问题等。设置环境变量JAX_SKIP_CUDA_CONSTRAINTS_CHECK可以关闭该版本检查见 jax_plugins/cuda/init.py但使用过旧的 CUDA 版本可能导致运行错误或结果不正确因此只建议在明确了解后果时使用。NCCL 是可选依赖仅在做多 GPU 计算时才需要。此外jax_plugins/cuda/init.py 的_load_nvidia_libraries展示了库加载策略优先从nvidia.*Python 包中加载 CUDA 运行时库libcudart.so、libcublas.so、libcudnn.so.9、libnccl.so.2等找不到时才回退到LD_LIBRARY_PATH——这从源码层面印证了不要乱设LD_LIBRARY_PATH的建议。5.5 NVIDIA GPU Docker 容器NVIDIA 提供了 JAX Toolbox 容器这些是包含 JAX 每夜构建版本以及部分模型/框架的前沿bleeding edge容器适合希望开箱即用并跟随最新代码的用户。六、Google Cloud TPU 安装JAX 为 Google Cloud TPU 提供预编译 wheel。在 Cloud TPU VM 中运行pip install jax[tpu]该命令会安装合适版本的jaxlib与libtpusetup.py 中tpuextra 还包含requests供jax.distributed.initialize使用。Colab 用户请务必使用TPU v2运行时而非已弃用的旧 TPU runtime。七、Mac GPUJAX不支持 macOS/OSX GPU在 Mac 上请直接使用上文第四节的标准 CPU 安装命令。Apple GPU 在支持平台表中为实验性档但官方安装文档明确指出 JAX 官方不支持 Mac/OSX GPU应使用 CPU 安装。八、AMD GPU 安装ROCmLinux8.1 ROCm 版本兼容性AMD GPU 支持由 AMD 维护的 ROCm JAX 插件提供。ROCm 官方兼容矩阵列出了受支持的 GPU SKU安装 JAX 之前请先按 AMD 的 ROCm 安装指南在系统上装好 ROCm。关键点JAX 目前不提供顺带安装 ROCm 本体的 pip extra。ROCm 必须已存在于宿主系统或容器内jax[rocm7-local]这个 extra 只负责在已有 ROCm 之上安装 JAX 的 ROCm 插件/PJRT 包。每个 JAX ROCm 插件版本针对特定 ROCm 版本构建因此已安装的 ROCm 必须与插件构建所针对的版本匹配。AMD 维护着权威的 JAX on ROCm compatibility matrix安装前请务必查阅确认目标 JAX 版本所需的 ROCm 版本rocm7插件包要求 ROCm 7.x 安装。这一约束在 setup.py 中体现为rocm7-local固定jax-rocm7-plugin{jax_version}.*。8.2 pip 安装预装 ROCm 路径推荐在已装好兼容 ROCm 的前提下pip install --upgrade jax[rocm7-local]ROCm 专属修复以插件的 post-release形式发布例如jax-rocm7-pluginX.Y.Z.post1。升级jax[rocm7-local]时会自动拾取配置的包索引中可用的最新兼容 post-release。验证安装python3 -c import jax; print(jax.devices())如果安装正常会列出你的 ROCm 设备例如[RocmDevice(id0), RocmDevice(id1), ...]。从源码层面看jax_plugins/rocm/init.py 的initialize()会先探测可用 AMD GPU 数量当检测到多 GPU 但/dev/shm不足 64 MB 时会发出警告——RCCL 在多 GPU 操作时可能耗尽共享内存导致运行时失败建议在 Docker 中使用--shm-size64g等参数增大/dev/shm。插件随后以platform_name ROCM注册 PJRT 插件。8.3 pip 安装ROCm wheel 直接安装技术预览AMD 正在推进从 AMD 自己的包索引直接安装 ROCm wheel目前是ROCm 7.13.0 中的技术预览尚未普遍可用。要点ROCm Core SDK 仍须单独安装JAX 包不会自动拉入rocm[libraries]各架构的索引 URL 与精确命令以 ROCm 7.13.0 安装指南为准作为同一预览工作的一部分ROCm 的 JAX 分支发布基于 TheRock ROCm 构建的 JAX wheel作为可下载的 release assets下载wheelhouse_*_theRock*.zip归档解压后pip install其中的jax-rocm7-pjrt与jax-rocm7-pluginwheel各版本的精确命令见对应 release notes。此路径同样是预览性质供评估使用。对于已正式发布的版本请使用上文 8.2 节的预装 ROCm 路径。8.4 AMD GPU Docker 容器AMD 提供了预构建的 ROCm JAX Docker 镜像打包了 ROCm、JAX 及全部依赖——这是最简单的起步方式因为无需自行安装 ROCmdocker pull rocm/jax:latest推荐的docker run参数与固定版本镜像标签请参考 AMD 官方的 JAX on ROCm 安装指南。Windows WSL2 注意ROCm 在 Windows WSL2 上的支持是实验性的。WSL 用户可能需要按 AMD 官方指南安装 ROCm for WSL在 WSL 环境中按标准 Linux ROCm JAX 安装步骤操作注意性能与稳定性可能与原生 Linux 安装存在差异。8.5 获取帮助ROCm JAX 插件由 AMD 维护。遇到 ROCm 专属问题安装、ROCm/驱动兼容性、插件或运行时错误请向 AMD 监控的 ROCm issue tracker 反馈非 ROCm 专属的 JAX 问题则反馈到 JAX issue tracker。九、Intel GPU 安装实验性Intel 提供了实验性的 OneAPI 插件intel-extension-for-openxla用于 Intel GPU 硬件。两种方式pip 安装参考 Intel 官方 JAX acceleration on Intel GPU 文档使用 Intel 的 XLA Docker 容器。问题反馈渠道JAX 相关问题反馈到 JAX issue trackerIntel OpenXLA 插件相关问题反馈到 intel-extension-for-openxla 的 issue tracker。从仓库源码看jax_plugins/oneapi/init.py 的initialize()会先尝试加载 Intel 的 SYCL/MKL 运行时库libsycl.so、libmkl_*、libur_adapter_level_zero.so等再以platform_name SYCL注册 PJRT 插件。jax[oneapi]extra 在 setup.py 中定义为安装jax-oneapi-plugin[with-oneapi]。注意该路径为实验性支持。十、Conda 安装社区支持conda-forge 提供了社区维护的jax构建conda install jax -c conda-forge在带 NVIDIA GPU 的机器上运行会安装 CUDA 版的jaxlib。若要确保装到的确实是 CUDA 版conda install jaxlib**cuda* jax -c conda-forge如需覆盖 JAX 使用的 CUDA 版本或在无 GPU 机器上安装 CUDA 构建可参考 conda-forge 网站 Tips tricks 中关于安装 CUDA 支持包如 TensorFlow、PyTorch 这类的说明更多细节见 conda-forge 的jaxlib与jaxfeedstock 仓库。十一、JAX Nightly 安装Nightly 版本反映构建时 JAX 主仓库的状态可能未通过完整测试套件。与正式版安装不同这里在命令行中显式列出 JAX 的所有包名以便 pip 在有新版本时逐一升级。JAX 将 nightly、RC候选版本与正式发布发布到多个非 PyPI 的 PEP 503 索引。所有 JAX 包都可以从统一索引https://us-python.pkg.dev/ml-oss-artifacts-published/jax/simple/获取该索引也镜像 PyPI 包这使得 nightly 安装可以使用 pip 的--index (-i)方式。注意统一索引在正式发布后的短时间内即使使用--pre也可能把 RC 或正式版当作最新版本此时最新 nightly 尚未重建。如果自动化或测试必须针对 nightly 进行或无法使用完整索引请使用只包含 nightly 工件的额外索引https://us-python.pkg.dev/ml-oss-artifacts-published/jax-public-nightly-artifacts-registry/simple/。各场景的 nightly 安装命令CPU onlypip install -U --pre jax jaxlib -i https://us-python.pkg.dev/ml-oss-artifacts-published/jax/simple/Google Cloud TPUpip install -U --pre jax jaxlib libtpu requests -i https://us-python.pkg.dev/ml-oss-artifacts-published/jax/simple/ -f https://storage.googleapis.com/jax-releases/libtpu_releases.htmlNVIDIA GPUCUDA 13pip install -U --pre jax jaxlib jax-cuda13-plugin[with-cuda] jax-cuda13-pjrt -i https://us-python.pkg.dev/ml-oss-artifacts-published/jax/simple/NVIDIA GPUCUDA 12pip install -U --pre jax jaxlib jax-cuda12-plugin[with-cuda] jax-cuda12-pjrt -i https://us-python.pkg.dev/ml-oss-artifacts-published/jax/simple/这些索引 URL 都是可直接浏览的 PEP 503 simple 仓库每个包有独立子目录例如jax包、jax-cuda12-pjrt包、jax-cuda13-pjrt包分别位于对应索引的独立路径下便于手工查看可用的工件列表。十二、从源码构建如果平台不在官方预编译 wheel 支持范围内或需要自定义构建配置请参考 docs/building_on_jax.md其中详细说明了从源码构建 JAX 与jaxlib的完整流程Bazel 构建、wheel 打包等。仓库根目录的 setup.py 与 jax/version.py 展示了版本号如何由构建环境变量JAX_RELEASE、JAX_NIGHTLY、WHEEL_VERSION_SUFFIX、JAX_GIT_HASH等控制从源码构建时可据此理解版本字符串的生成规则。十三、安装旧版 jaxlib wheels受 PyPI 存储空间限制JAX 团队会定期从pypi.org/project/jax的 releases 中移除较旧的jaxlibwheel。这些旧 wheel 仍可通过如下 URL 直接安装# 通过 wheel 归档在 CPU 上安装 jax pip install jax[cpu]0.3.25 -i https://us-python.pkg.dev/ml-oss-artifacts-published/jax/simple/ # 直接安装 jaxlib 0.3.25 CPU wheel pip install jaxlib0.3.25 -i https://us-python.pkg.dev/ml-oss-artifacts-published/jax/simple/特定旧版 GPU wheel 需要使用jax_cuda_releases.htmlURL例如pip install jaxlib0.3.25cuda11.cudnn82 -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html十四、安装后的验证与常见问题速查无论选择哪种平台安装完成后都建议先验证环境python3 -c import jax; print(jax.devices()) python3 -c import jax; print(jax.__version__)常见问题的排查要点汇总jaxlib未安装 / 版本不匹配表现为导入jax报ModuleNotFoundError或RuntimeError。原因是jax与jaxlib版本未对齐见 jax/_src/lib/init.py 的校验逻辑。解决使用pip install -U jax[extra]完整重装让 pip 自动解析匹配版本。CUDA 组件版本过旧报 Outdated ... installation found附带 Version JAX was built against / Minimum supported / Installed version 信息。解决升级对应 NVIDIA 库或按 5.2 节重装 wheel确有必要时可用JAX_SKIP_CUDA_CONSTRAINTS_CHECK跳过检查但可能引发错误结果。LD_LIBRARY_PATH干扰导致加载到错误版本的 CUDA 库表现为启动即崩溃或版本检查失败。解决清空LD_LIBRARY_PATH后重试。多 GPU 环境 ROCm 失败注意/dev/shm容量是否充足参考 jax_plugins/rocm/init.py 的警告逻辑容器场景适当调大--shm-size。不支持平台的静默失败Windows 上使用-localGPU wheel 或非支持架构会看似成功安装但运行失败务必先对照第三节的支持矩阵。以上安装路径均可在仓库 docs/installation.md 中找到原文依据相关版本约束与插件依赖可在 setup.py、jax/version.py 及各 jax_plugins 子目录的插件实现中进一步核实。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考