资讯动态

从源码构建 JAX:jaxlib、hermetic Python、测试与文档开发全指南

发布时间:2026/9/20 21:09:03 来源:尧图企业网站定制
从源码构建 JAXjaxlib、hermetic Python、测试与文档开发全指南【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax本文是 JAX当前仓库官方开发者文档 docs/developer.md 的深度展开版系统讲解从源码构建 JAX 的完整流程如何获取源码、构建或安装jaxlib含 CUDA/ROCM/Windows/XLA 定制如何利用 hermetic Python 机制锁定可复现的构建环境如何通过 Bazel 或 pytest 运行测试以及如何做类型检查、Lint 与文档维护。读完本文你将能够在自己机器上搭建完整的 JAX 开发环境、构建多平台多后端的 wheel并熟练运用仓库内提供的构建与测试工具链。JAX 的构建模型jaxlib 与 jax 的双包结构JAX 的源码构建本质上分为两个独立步骤这是理解整篇文档的前提构建或安装jaxlib这是 JAX 的 C 支持库内含 XLA 编译器、运行时与各类内核实现对应仓库中的 jaxlib/ 目录。安装jaxPython 包即仓库根目录下的纯 Python 包包含 jax/ 目录中的全部 Python 代码。两者的安装是解耦的如果你只修改 JAX 的 Python 部分完全可以跳过 C 编译直接用 pip 装一个预编译好的jaxlibwheel再把jax以可编辑模式pip install -e装进环境实现改 Python 代码即时生效、C 部分用官方二进制的开发体验。获取源码的方式很直接文档中约定以python作为 Python 3 解释器名称部分系统需改用python3git clone https://github.com/google/jax cd jax本文中出现的所有build.py调用均指仓库根目录下build/build.py它是整个构建流程的入口脚本见下文源码剖析。构建或安装 jaxlib方式一用 pip 安装预编译 jaxlib仅改 Python 代码时推荐如果开发工作只涉及 Python 层例如 jax/_src/numpy/lax_numpy.py 这类纯 Python 实现文档建议直接从 PyPI 安装预编译 wheelpip install jaxlibGPU 与 TPU 支持所需的额外配置CUDA wheel 选择、TPU 版本等请参考 README.md 中关于 pip 安装的完整说明。方式二从源码构建 jaxlib从源码构建需要先准备 C 编译工具链LinuxUbuntu/Debiansudo apt install g python python3-devmacOS安装 Xcode 及 Xcode Command Line Tools。Windows见下文专节。一个关键设计是构建过程不需要本地安装 Python 依赖——Bazel 会使用自己的 hermetic Python 解释器详见Managing hermetic Python小节你系统里的 Python 在构建期间会被忽略只有build/build.py脚本本身由系统 Python 解释。CPU 或 TPU 目标的标准构建命令python build/build.py pip install dist/*.whl # 安装 jaxlib内含 XLA默认情况下 wheel 会输出到当前目录下的dist/子目录。要为与系统不同的 Python 版本构建 wheel追加--python_version参数python build/build.py --python_version3.12build.py还支持大量配置选项运行python build/build.py --help可查看全部参数包括指定 CUDA/CuDNN 路径的方式。无论是否传--python_versionBazel 侧始终使用 hermetic Python。CUDA 支持两种构建策略文档给出了两条构建 CUDA 版 jaxlib 的路径单 wheel 方案——把 CUDA 支持直接打进 jaxlib wheelpython build/build.py --enable_cudaGPU 插件分离方案——生成三个 wheel不含 CUDA 的 jaxlib、jax-cuda-plugin、jax-cuda-pjrtpython build/build.py --enable_cuda --build_gpu_plugin --gpu_plugin_cuda_version12gpu_plugin_cuda_version可设为 11 或 12。从源码看build/build.py在检测到--enable_cuda后会写入.jax_configure.bazelrc其中包含build --configcuda、TF_CUDA_PATHS、TF_CUDA_VERSION、TF_CUDNN_VERSION、TF_CUDA_COMPUTE_CAPABILITIES等 Bazel action env--build_gpu_plugin则会追加build --configcuda_pluginbuild/build.py。由此可见--enable_cuda本质上是把 CUDA 工具链通过 Bazel 配置注入到 XLA 的构建中--build_gpu_plugin则切换到独立的 GPU 插件构建目标。使用本地修改版 XLA 仓库构建JAX 依赖 XLA其源码位于独立的 XLA 仓库默认使用 JAX 固定的pinnedXLA 提交。当你在本地修改 XLA 时有两种方式让构建使用你的副本Bazeloverride_repository特性以命令行参数形式传给build.pypython build/build.py --bazel_options--override_repositoryxla/path/to/xla直接修改 JAX 源码树根部的 WORKSPACE将 XLA 指向不同的 tree。若要回馈 XLA 改动请向 XLA 仓库提交 PR。JAX 所固定的 XLA 版本会定期更新尤其是每次jaxlib发布之前。Windows 上的 jaxlib 构建要点Windows 构建需要额外满足以下条件C 工具链安装 Visual Studio2019 16.5 或更新版本若需 CUDA按 NVIDIA CUDA 安装指南配置 CUDA 环境。符号链接JAX 构建使用符号链接因此必须开启 Windows 的Developer Mode。Python 环境可使用 Python 官方 Windows 安装包或 Anaconda / Miniconda。MSYS2Bazel 的部分目标依赖 bash 工具做脚本化需安装 MSYS2并安装patch与coreutils后者提供realpath命令pacman -S patch coreutils一切就绪后在 PowerShell 中确保bazel、patch、realpath可访问激活 conda 环境后即可构建以下示例开启了 CUDA可按需调整路径与版本python .\build\build.py --enable_cuda --cuda_pathC:/Program Files/NVIDIA GPU Computing Toolkit/CUDA/v10.1 --cudnn_pathC:/Program Files/NVIDIA GPU Computing Toolkit/CUDA/v10.1 --cuda_version10.1 --cudnn_version7.6.5如需带调试信息构建追加--bazel_options--copt/Z7。为 AMD GPU 构建 ROCM 版 jaxlib构建 ROCM 版需要若干 ROCM/HIP 库。以配置了 AMDapt仓库的 Ubuntu 为例sudo apt install miopen-hip hipfft-dev rocrand-dev hipsparse-dev hipsolver-dev \ rccl-dev rccl hip-dev rocfft-dev roctracer-dev hipblas-dev rocm-device-libs构建命令按实际路径与 ROCM 版本调整python build/build.py --enable_rocm --rocm_path/opt/rocm-5.7.0AMD 的 XLA fork 可能包含上游没有的修复。若上游 XLA 遇到问题可克隆 AMD fork 并覆盖构建所用的 XLA 仓库git clone https://github.com/ROCmSoftwarePlatform/xla.git python build/build.py --enable_rocm --rocm_path/opt/rocm-5.7.0 \ --bazel_options--override_repositoryxla/path/to/xla-rocm与 CUDA 路径一致--enable_rocm也会在.jax_configure.bazelrc中写入build --configrocm与ROCM_PATH、TF_ROCM_AMDGPU_TARGETS等配置build/build.py。build/rocm/build_rocm.sh中还提供了针对 ROCM 环境的封装脚本可供参考。管理 hermetic Python可复现构建的核心机制JAX 的所有 Bazel 构建与测试命令都依赖hermetic Python基于 rules_pythonBazel 自行管理 Python 解释器及其全部依赖从而保证构建在 Linux、Windows、macOS 上行为一致且与本地系统的 Python 完全隔离。指定 Python 版本运行build/build.py时hermetic Python 的版本会自动匹配你用来执行脚本的 Python 版本也可以用--python_version显式指定python build/build.py --python_version3.12底层由HERMETIC_PYTHON_VERSION环境变量控制。build/build.py会自动设置它若直接运行 bazel则需要手动设置三种方式任选# 方式一写入 .bazelrc build --repo_envHERMETIC_PYTHON_VERSION3.12 # 方式二直接在构建命令中传递 bazel build target --repo_envHERMETIC_PYTHON_VERSION3.12 # 方式三在 shell 中全局导出 export HERMETIC_PYTHON_VERSION3.12值得注意的是build.py本身也会把HERMETIC_PYTHON_VERSION写入生成的.jax_configure.bazelrcbuild/build.py这意味着用build.py构建时版本一致性是自动保证的无需手动干预。因为不同 Python 版本共享与解释器无关的构建缓存你可以在同一台机器上通过切换--python_version依次对不同版本的 Python 构建和测试之前的缓存会被保留复用。指定 Python 依赖requirements 锁定文件为保证构建可复现Bazel 构建期间 JAX 的全部 Python 依赖都被钉死到特定版本。完整的依赖传递闭包及其哈希记录在build/requirements_lock_python version.txt文件中例如 Python 3.12 对应 build/requirements_lock_3_12.txt当前仓库提供了 3.103.13 的锁定文件。直接依赖清单维护在 build/requirements.in 中其中既有运行时依赖numpy、scipy、ml_dtypes、opt_einsum、zstandard、etils[epath]也有测试依赖通过-r test-requirements.txt引入。更新锁定文件的命令python build/build.py --requirements_update --python_version3.12该命令底层调用pip-compilepip-tools。如果希望有更多控制也可以直接运行等价的 Bazel 命令bazel run //build:requirements.update --repo_envHERMETIC_PYTHON_VERSION3.12其中3.12是你希望更新的 Python 版本。由于底层仍是pip与pip-compile这两个工具支持的绝大多数命令行参数都会被接受例如想让更新器考虑预发布版本bazel run //build:requirements.update --repo_envHERMETIC_PYTHON_VERSION3.12 -- --pre依赖本地 wheel如果需要依赖本地.whl文件比如你刚构建好的 jaxlib wheel把 wheel 的路径追加到 build/requirements.in 并重新运行更新器即可echo -e \n$(realpath jaxlib-0.4.27.dev20240416-cp312-cp312-manylinux2014_x86_64.whl) build/requirements.in python build/build.py --requirements_update --python_version3.12依赖 nightly wheel要针对最新、可能不稳定的一组依赖构建测试使用 nightly 版本更新器python build/build.py --requirements_nightly_update --python_version3.12等价 Bazel 命令bazel run //build:requirements_nightly.update --repo_envHERMETIC_PYTHON_VERSION3.12与常规更新器的区别在于它默认接受预发布、dev 与 nightly 包会额外搜索pypi.anaconda.org/scientific-python-nightly-wheels作为 extra index url且最终锁定文件不写入哈希。用预发布版 Python 构建进阶JAX 开箱即支持所有当前正式发布的 Python 版本若要针对尚未正式发布的版本构建测试按以下步骤安装编译 Python 解释器及其关键包如 numpy/scipy所需的系统包。典型 Debian 系统sudo apt-get update sudo apt-get build-dep python3 -y sudo apt-get install pkg-config zlib1g-dev libssl-dev -y # 构建 scipy 需要 sudo apt-get install libopenblas-dev -y检查 WORKSPACE确认其中存在指向目标 Python 版本的custom_python_interpreter()条目。构建 Python 解释器bazel build python_dev//:python_dev默认用 GCC 编译若想用 clang可设置对应环境变量例如--repo_envCC/usr/lib/llvm-17/bin/clang --repo_envCXX/usr/lib/llvm-17/bin/clang。把生成的python_register_toolchains()片段写入 WORKSPACE。上一步命令的末尾会打印一段代码片段将其复制到python_init_toolchains()条目之后新增版本或替换之如用自建 3.12 替换默认 3.12。片段已按你的实际环境生成可直接使用也可自定义比如把 Python 的.tgz改为远程下载地址。确保 WORKSPACE 的python_init_repositories()的requirements参数包含该版本条目例如 Python 3.13 应有3.13: //build:requirements_lock_3_13.txt。推荐对不稳定版本预构建全部 Python 依赖bazel build //build:all_py_deps --repo_envHERMETIC_PYTHON_VERSION3.13这会让 pip 对尚无二进制分发的包如 numpy、scipy、matplotlib、zstandard从源码拉取构建。建议在实际 JAX 构建之前独立执行此步骤以免两者互相冲突。例如 JAX 通常用 clang 构建而matplotlib从源码构建时假设使用 GCCclang 会因 LTO-flto行为差异导致构建失败。若针对稳定版本 Python 构建、或所有依赖都有现成二进制分发可跳过此步。构建完成后正常执行构建/测试命令即可只需确保HERMETIC_PYTHON_VERSION指向你的新版本。关于锁定文件更新的提醒对预发布版本 Python 直接更新requirements_lock_python_version.txt很可能失败——仓库里没有匹配的二进制包时pip-compile会尝试从源码构建这比pip安装更严格更容易失败。推荐做法是为最新稳定版本如 3.12生成不带哈希的锁定文件再复制给不稳定版本如 3.13bazel run //build:requirements_dev.update --repo_envHERMETIC_PYTHON_VERSION3.12 cp build/requirements_lock_3_12.txt build/requirements_lock_3_13.txt bazel build //build:all_py_deps --repo_envHERMETIC_PYTHON_VERSION3.13 # 根据依赖对新版本 Python 的兼容程度你可能需要手动编辑最终的锁定文件安装 jax Python 包jaxlib就绪后在仓库根部用可编辑模式安装jaxpip install -e . # 安装 jax要从 GitHub 升级到最新版只需在仓库根部执行git pull然后按需重新运行build.py或升级jaxlib。通常不需要重装jax——pip install -e已经在 site-packages 和仓库之间建立了符号链接。运行测试JAX 测试支持 Bazel 与 pytest 两种机制。使用 Bazel 运行测试首先配置构建python build/build.py --configure_only可按需向build.py追加其他配置选项。默认情况下 Bazel 测试使用从源码构建的 jaxlib运行方式bazel test //tests:cpu_tests //tests:backend_independent_tests//tests:gpu_tests与//tests:tpu_tests目标同样可用需要相应硬件。如果希望改用预装的 jaxlib而非每次重新构建需要先把 jaxlib 装进 hermetic Python。安装指定版本以jaxlib 0.4.26为例echo -e \njaxlib 0.4.26 build/requirements.in python build/build.py --requirements_update或从本地 wheel 安装假设 Python 3.12echo -e \n$(realpath jaxlib-0.4.26-cp312-cp312-manylinux2014_x86_64.whl) build/requirements.in python build/build.py --requirements_update --python_version3.12hermetic 环境装好 jaxlib 后用如下命令跑测试bazel test --//jax:build_jaxlibfalse //tests:cpu_tests //tests:backend_independent_tests多加速器测试部分测试面向多 GPU/TPU。JAX 已安装时可用如下方式运行 GPU 测试bazel test //tests:gpu_tests --local_test_jobs4 --test_tag_filtersmultiaccelerator --//jax:build_jaxlibfalse --test_envXLA_PYTHON_CLIENT_ALLOCATORplatform单加速器测试并行加速可在多个加速器上并行运行每个加速器同时跑多个测试。以 2 块 GPU、每块 4 个并发任务为例NB_GPUS2 JOBS_PER_ACC4 J$((NB_GPUS * JOBS_PER_ACC)) MULTI_GPU--run_under $PWD/build/parallel_accelerator_execute.sh --test_envJAX_ACCELERATOR_COUNT${NB_GPUS} --test_envJAX_TESTS_PER_ACCELERATOR${JOBS_PER_ACC} --local_test_jobs$J bazel test //tests:gpu_tests //tests:backend_independent_tests --test_envXLA_PYTHON_CLIENT_PREALLOCATEfalse --test_tag_filters-multiaccelerator $MULTI_GPU这里的 build/parallel_accelerator_execute.sh 就是仓库中提供的多加速器执行包装脚本通过JAX_ACCELERATOR_COUNT与JAX_TESTS_PER_ACCELERATOR两个环境变量控制资源划分。使用 pytest 运行测试先安装依赖pip install -r build/test-requirements.txtbuild/test-requirements.txt 中已包含pytest-xdist因此可以从仓库根部直接并行运行全部测试pytest -n auto tests-n auto由 pytest-xdist 提供会按 CPU 核数自动并行。控制测试行为JAX 会组合式地生成测试用例每个测试默认生成并校验 10 个用例可用JAX_NUM_GENERATED_CASES环境变量调整CI 自动化测试默认使用 25# Bazel bazel test //tests/... --test_envJAX_NUM_GENERATED_CASES25 # pytest JAX_NUM_GENERATED_CASES25 pytest -n auto tests自动化测试还会以 64 位浮点/整数模式运行JAX_ENABLE_X64JAX_ENABLE_X641 JAX_NUM_GENERATED_CASES25 pytest -n auto tests运行单个测试文件可看到更详细的用例信息JAX_NUM_GENERATED_CASES5 python tests/lax_numpy_test.py跳过已知慢测试JAX_SKIP_SLOW_TESTS1。用--test_targets指定文件内的特定测试支持字符串或正则例如运行jax.numpy.pad的全部测试python tests/lax_numpy_test.py --test_targetstestPad文档构建流程中还会校验 Colab notebook 无报错。DoctestsJAX 用 pytest 的 doctest 模式测试文档中的代码示例pytest docs此外还以doctest-modules模式确保函数 docstring 中的示例可运行例如pytest --doctest-modules jax/_src/numpy/lax_numpy.py注意对整个包执行 doctest 时有若干文件被标记为跳过细节可在 CI 工作流配置仓库中的.github/workflows/ci-build.yaml中查看。类型检查JAX 使用mypy检查类型标注与 CI 完全一致的本机检查命令pip install mypy mypy --configpyproject.toml --show-error-codes jaxmypy 配置位于仓库根部的 pyproject.toml。也可以借助 .pre-commit-config.yaml 中定义的 pre-commit 钩子对 git 暂存区文件自动运行与 GitHub CI 相同版本的 mypypre-commit run mypyLintingJAX 使用ruff保证代码质量本机检查pip install ruff ruff jax同样支持通过 pre-commit 对暂存文件自动执行与 CI 相同版本pre-commit run ruff更新文档重建文档Sphinx安装文档构建依赖pip install -r docs/requirements.txt然后构建 HTML 文档sphinx-build -b html docs docs/build/html -j auto由于构建会执行文档源中的大量 notebook耗时可能很长跳过 notebook 执行可加速sphinx-build -b html -D nb_execution_modeoff docs docs/build/html -j auto生成结果位于docs/build/html/index.html。-j auto控制构建并行度也可换成具体数字以指定 CPU 核心数。维护 notebookipynb 与 md 双格式文档使用 jupytext 中维护两份同步的 notebookipynb格式可直接在 Colab 打开执行md格式则便于版本控制中查看 diff仓库中如 docs/notebooks/thinking_in_jax.md 与同名.ipynb即为一对。编辑 ipynb改动较大涉及代码与输出时建议在 Jupyter 或 Colab 中编辑完成后Run all cells再Download ipynb并按上文sphinx-build方式验证可执行。编辑 md仅修改文本内容时直接编辑.md版本更方便。同步双版本编辑完任一版本后用 jupytext 同步版本应与 .pre-commit-config.yaml 中指定的一致pip install jupytext1.16.0 jupytext --sync docs/notebooks/thinking_in_jax.ipynb校验同步是否正确的 pre-commit 钩子git add docs -u # pre-commit 只处理暂存区文件 pre-commit run jupytext新建 notebook若要纳入 jupytext 同步先设置格式会在 notebook 中写入jupytext元数据字段jupytext --sync据此识别jupytext --set-formats ipynb,md:myst path/to/the/notebook.ipynbSphinx 构建中的 notebook部分 notebook 会在预提交检查与 Read the Docs 构建中自动执行单元格报错会导致构建失败。如果报错是有意为之可在.ipynb中手动为单元格打上raises-exceptions元数据保存时会被保留。含长计算等场景的 notebook 会通过 docs/conf.py 的exclude_patterns排除出构建。Read the Docs 上的文档构建JAX 的自动化文档位于 jax.readthedocs.io构建由项目级 Read the Docs 设置驱动代码推送到main分支即触发文档构建。每次构建由 .readthedocs.yml 与 docs/conf.py 驱动——前者定义了构建镜像ubuntu-22.04、Python 版本3.10、Sphinx 配置路径docs/conf.py并开启fail_on_warning: true以及额外格式htmlzip和docs/requirements.txt依赖安装。推送到test-docs分支也会自动构建可用于预演文档效果。本地复现 Read the Docs 构建可参考其构建日志在全新目录中依次执行依赖安装与sphinx-build命令例如创建虚拟环境、git clone --no-single-branch --depth 50拉取仓库、切换test-docs分支、安装docs/requirements.txt最后运行python \which sphinx-build -T -E -b html -d _build/doctrees-readthedocs -D languageen . _build/html。附build.py 底层实现速览整个构建体系的入口是 build/build.py共 749 行从源码可以确认几个关键实现细节Bazel 自动获取与校验脚本内置了各平台Linux x86_64/aarch64、Darwin x86_64/arm64、Windows AMD64的 Bazel 6.5.0 下载映射找不到本地 bazel 时自动下载并用内置 SHA256 校验二进制完整性build/build.py。Bazel 版本下限get_bazel_path要求 bazel 版本 6.5.0否则报错退出build/build.py。Python 版本下限check_python_version要求 Python 3.10 或更新版本build/build.py这与仓库提供requirements_lock_3_10.txt~requirements_lock_3_13.txt的覆盖范围一致。配置生成write_bazelrc根据--enable_cuda、--enable_rocm、--build_gpu_plugin、--use_clang、--target_cpu_features、--python_version等参数生成.jax_configure.bazelrc把 CUDA/ROCM 路径、编译配置与HERMETIC_PYTHON_VERSION注入 Bazelbuild/build.py。CPU 特性选项--target_cpu_features支持releasex86-64 下启用 AVX、native-marchnativeWindows 不支持、default交给编译器默认行为build/build.py。理解这些底层机制有助于你在自定义构建参数或排查构建问题时定位到build.py的具体分支逻辑。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价