如何本地构建带[cuda12]标签的jax?已成功编译jaxlib
我已成功编译jaxlib,但不知如何构建带[cuda12]标签的jax?在jax文档中未找到构建该特定wheel的相关说明,恳请解答!
我有部分应用运行于通过pip install jax[cuda12]==0.4.34 jaxlib==0.4.34安装的jax环境。近期同时启用cuDNN与持久缓存时,遇到错误:String field 'xla.gpu.CompilationResultProto.DnnCompiledGraphsEntry.value' contains invalid UTF-8 data。
我发现该bug已在2025年1月9日xla仓库的某次提交中修复,对应jax最新版本为0.5.0。于是卸载旧版本后通过pip install jax[cuda12]==0.5.0 jaxlib==0.5.0重新安装,上述问题得以解决!
但我的应用与jax 0.5.0不兼容,运行速度变慢且出现NaN错误,因此决定基于本地xla仓库构建jaxlib与jax[cuda12]。
我拉取了jax与xla仓库并切换到特定提交,命令如下:
git clone --recurse-submodules https://github.com/jax-ml/jax.git git clone --recurse-submodules https://github.com/openxla/xla.git cd jax # 对应jax v0.4.34版本 git checkout affba367c5533df8900e32cbc3d31ca92dd1c1ea git submodule update --init --recursive cd .. cd xla # 对应jax v0.4.34版本中定义的XLA版本 git checkout cd6e808c59f53b40a99df1f1b860db9a3e598bff git submodule update --init --recursive
修改XLA源码修复bug后,我参考开发者文档,使用以下命令编译jaxlib:
python3 build/build.py \ --python_version=3.11 \ --enable_cuda \ --cuda_version=12.6.1 \ --cudnn_version=9.4.0 \ --bazel_options=--override_repository=xla=/home/thomas/xla \ --verbose
等待完成后得到如下输出:
Target //jaxlib/tools:build_wheel up-to-date: bazel-bin/jaxlib/tools/build_wheel INFO: Elapsed time: 1928.511s, Critical Path: 287.80s INFO: 3390 processes: 24 internal, 3366 local. INFO: Build completed successfully, 3390 total actions INFO: Running command line: bazel-bin/jaxlib/tools/build_wheel '--output_path=/home/thomas/jax/dist' '--jaxlib_git_hash=affba367c5533df8900e32cbc3d31ca92dd1c1ea' '--cpu=x86_64' Output wheel: /home/thomas/jax/dist/jaxlib-0.4.34.dev20250218-cp311-cp311-manylinux2014_x86_64.whl To install the newly-built jaxlib wheel on system Python, run: pip install /home/thomas/jax/dist/jaxlib-0.4.34.dev20250218-cp311-cp311-manylinux2014_x86_64.whl --force-reinstall To install the newly-built jaxlib wheel on hermetic Python, run: echo -e "\n/home/thomas/jax/dist/jaxlib-0.4.34.dev20250218-cp311-cp311-manylinux2014_x86_64.whl" >> build/requirements.in bazel run //build:requirements.update --repo_env=HERMETIC_PYTHON_VERSION=3.11
首先明确:jax[cuda12]本质是pip的额外依赖组,并非特殊的wheel包。官方发布的jax wheel本身就包含所有额外依赖组的声明,你只需要构建标准的jax wheel,然后安装时指定[cuda12]即可,同时配合自己编译的jaxlib。
1. 构建标准jax wheel
在已切换到指定提交的jax仓库根目录下,执行以下命令构建jax的wheel包:
# 安装构建依赖 pip install build # 构建wheel python -m build --wheel
执行完成后,dist目录下会生成类似jax-0.4.34.dev20250218-cp311-cp311-manylinux_2_17_x86_64.whl的文件。
2. 安装带[cuda12]依赖的jax(配合自定义jaxlib)
因为你已经编译了适配CUDA 12的jaxlib,安装时可按以下步骤操作:
# 先安装自己编译的jaxlib pip install /home/thomas/jax/dist/jaxlib-0.4.34.dev20250218-cp311-cp311-manylinux2014_x86_64.whl --force-reinstall # 再安装jax并启用cuda12依赖组 pip install /path/to/your/jax/dist/jax-0.4.34.dev20250218-cp311-cp311-manylinux_2_17_x86_64.whl[cuda12]
也可以合并为一条命令(确保路径正确):
pip install "jax[cuda12] @ file:///path/to/your/jax/dist/jax-0.4.34.dev20250218-cp311-cp311-manylinux_2_17_x86_64.whl" \ --force-reinstall \ --find-links=/home/thomas/jax/dist
--find-links参数会让pip优先从指定目录查找jaxlib,而非从PyPI下载官方版本。
3. (可选)自定义cuda12依赖版本
如果你想让jax[cuda12]直接指向你编译的jaxlib版本,可以修改jax仓库根目录下的setup.py文件,找到extras_require中的cuda12条目,将其中的jaxlib版本替换为你编译的版本(比如jaxlib==0.4.34.dev20250218),修改后示例:
extras_require={ # ...其他依赖组 "cuda12": [ "jaxlib==0.4.34.dev20250218", "nvidia-cuda-runtime-cu12>=12.0", "nvidia-cudnn-cu12>=8.9", # ...其他依赖 ], # ... }
修改完成后重新执行构建命令即可。
补充说明
jax[cuda12]的作用是自动安装适配CUDA 12的jaxlib及NVIDIA相关依赖(如cuda-runtime、cudnn等),无需专门构建带标签的jax wheel。- 构建jax wheel不需要依赖CUDA环境,因为jax本身是纯Python包,CUDA相关逻辑都在jaxlib中。
内容的提问来源于stack exchange,提问作者Thomas

