GPU运行简单JAX程序出现内存错误问题求助
问题背景
通过命令pip install --upgrade "jax[cuda11_pip]" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html安装JAX后,运行以下简单代码:
import jax.numpy as jnp a = jnp.array([1,2,3]) a.dot(a)
触发CUDNN初始化错误:
2023-09-08 10:12:55.791658: E external/xla/xla/stream_executor/cuda/cuda_dnn.cc:445] Could not create cudnn handle: CUDNN_STATUS_INTERNAL_ERROR 2023-09-08 10:12:55.791696: E external/xla/xla/stream_executor/cuda/cuda_dnn.cc:449] Memory usage: 8058437632 bytes free, 8513978368 bytes total.
系统nvidia-smi输出:
NVIDIA-SMI 470.199.02 Driver Version: 470.199.02 CUDA Version: 11.4 | |-------------------------------+----------------------+----------------------+ | GPU Name Persistence-M| Bus-Id Disp.A | Volatile Uncorr. ECC | | Fan Temp Perf Pwr:Usage/Cap| Memory-Usage | GPU-Util Compute M. | | | | MIG M. | |===============================+======================+======================| | 0 NVIDIA GeForce ... Off | 00000000:01:00.0 Off | N/A | | N/A 64C P0 37W / N/A | 364MiB / 8119MiB | 2% Default | | | | N/A | +-------------------------------+----------------------+----------------------+
已尝试JAX官方文档中的内存优化建议,但问题未解决。
解决思路
校验CUDNN与CUDA版本兼容性
当前系统CUDA版本为11.4,需确认JAX依赖的CUDNN版本是否与之匹配。版本不兼容会直接导致CUDNN初始化失败,可尝试安装对应CUDA 11.4的CUDNN版本,或调整CUDA驱动至兼容版本。强制CPU运行验证问题根源
设置环境变量export JAX_PLATFORM_NAME=cpu后重新运行代码,若正常执行,说明问题出在GPU/CUDNN组件,而非JAX核心逻辑。清理GPU残留进程
尽管nvidia-smi显示内存占用低,但僵尸进程可能占用隐性资源。用nvidia-smi --query-compute-apps=pid,process_name --format=csv排查关联进程,杀掉无关GPU进程后重试。绕过CUDNN初始化
在代码开头添加配置禁用CUDNN:import jax jax.config.update('jax_enable_xla', False)或设置环境变量
export XLA_FLAGS=--xla_gpu_cudnn_disable=true,验证是否能正常执行。重装适配版本的JAX
完全卸载现有JAX组件:pip uninstall -y jax jaxlib,然后安装适配CUDA 11.4的指定版本,例如:pip install jax jaxlib==0.4.14+cuda11.cudnn86 -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html注意替换版本号为对应CUDA 11.4的兼容版本。
检查GPU设备权限
确保当前用户已加入video或nvidia用户组,避免因权限不足导致CUDNN初始化失败。
内容的提问来源于stack exchange,提问作者tohe

