使用Apptainer在GPU上运行JAX时出现CuDNN错误
JAX在Apptainer容器中GPU运行失败的解决方案
问题描述
基于Python 3.10+和JAX开发的应用,需在配备NVIDIA A40 GPU的集群上运行,集群仅支持Apptainer(不支持Docker)。参考JAX官方Dockerfile制作Ubuntu镜像并转换为Apptainer镜像后,运行JAX代码时出现CuDNN初始化失败或库文件缺失的错误。
初始镜像配置
FROM nvidia/cuda:11.8.0-devel-ubuntu22.04 RUN apt update && apt install python3-pip -y RUN pip install "jax[cuda11_cudnn86]" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html
错误1:CuDNN初始化失败
运行代码:
import jax jax.numpy.array(1.)
报错:
2023-05-11 14:41:50.580441: E external/xla/xla/stream_executor/cuda/cuda_dnn.cc:429] Could not create cudnn handle: CUDNN_STATUS_INTERNAL_ERROR ... jaxlib.xla_extension.XlaRuntimeError: FAILED_PRECONDITION: DNN library initialization failed. Look at the errors above for more details.
错误2:库文件缺失
改用包含CuDNN 8.7.0的基础镜像后:
FROM nvidia/cuda:11.8.0-cudnn8-runtime-ubuntu22.04 RUN apt update && apt install python3-pip -y RUN pip install "jax[cuda11_cudnn86]" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html
运行时报错:
Could not load library libcudnn_ops_infer.so.8. Error: libnvrtc.so: cannot open shared object file: No such file or directory Aborted (core dumped)
已验证的前提条件
- 无GPU的本地Docker环境运行该镜像正常
- Apptainer容器可通过
nvidia-smi正常检测到NVIDIA A40 GPU - 手动添加CUDA路径后问题未解决
解决方案
1. 对齐宿主机与容器的CUDA版本
从nvidia-smi输出可知,宿主机CUDA驱动版本为11.7,而初始镜像使用的是CUDA 11.8。虽然驱动向前兼容,但版本差异可能引发CuDNN兼容性问题,建议使用与宿主机驱动匹配的CUDA镜像版本:
FROM nvidia/cuda:11.7.1-devel-ubuntu22.04
2. 匹配JAX依赖的CuDNN版本
JAX安装参数cuda11_cudnn86对应CuDNN 8.6.x,需使用包含该版本CuDNN的基础镜像,避免版本不匹配:
FROM nvidia/cuda:11.7.1-cudnn8.6-devel-ubuntu22.04
3. 固定JAX版本避免兼容性问题
明确指定JAX和jaxlib的版本,防止自动升级到与CUDA/CuDNN不兼容的版本:
FROM nvidia/cuda:11.7.1-cudnn8.6-devel-ubuntu22.04 RUN apt update && apt install python3-pip -y RUN pip install "jax==0.4.10" "jaxlib==0.4.10+cuda11.cudnn86" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html
4. 正确配置Apptainer运行环境
启动容器时,确保Apptainer正确传递GPU相关环境变量,可显式指定库路径:
apptainer run --nv --env LD_LIBRARY_PATH=/usr/local/cuda/lib64:/usr/lib/x86_64-linux-gnu docker://my-image
5. 验证GPU可用性
进入容器后运行以下代码确认JAX是否识别到GPU:
import jax print(jax.devices())
若输出包含类似[GpuDevice(id=0, process_index=0), ...]的内容,则说明配置成功。
内容的提问来源于stack exchange,提问作者Hylke
相关产品推荐
相关产品推荐

