You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

GPU运行简单JAX程序出现内存错误问题求助

JAX CUDA/CUDNN 内存错误排查与解决思路

问题背景

通过命令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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.11 16:20:17