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

使用JAX在Colab运行图像生成Notebook遇DeviceArrayBase属性错误

问题描述

在Colab运行基于JAX的图像生成Notebook时,遇到两个问题:

  1. GPU/TPU未被检测到,自动 fallback 到CPU
  2. 导入jaxtorch时触发AttributeError,提示module 'jaxlib.xla_extension' has no attribute 'DeviceArrayBase'

错误栈详情:

WARNING:jax._src.xla_bridge:No GPU/TPU found, falling back to CPU. (Set TF_CPP_MIN_LOG_LEVEL=0 and       rerun for more info.)

---------------------------------------------------------------------------

AttributeError                            Traceback (most recent call last)

<ipython-input-7-73b0723cc3af> in <cell line: 23>()
 21 import jax.numpy as jnp
 22 import jax.scipy as jsp
---> 23 import jaxtorch
 24 from jaxtorch import PRNG, Context, Module, nn, init
 25 from tqdm import tqdm

3 frames

/content/./jax-guided-diffusion/jaxtorch/monkeypatches.py in register(**kwargs)
 16             print(f'Not monkeypatching DeviceArray and Tracer with `{attr}`, because that method is already implemented.', file=sys.stderr)
 17             continue
---> 18         setattr(jaxlib.xla_extension.DeviceArrayBase, attr, fun)
 19         setattr(jax.interpreters.xla.DeviceArray, attr, fun)
 20         setattr(jax.core.Tracer, attr, fun)

AttributeError: module 'jaxlib.xla_extension' has no attribute 'DeviceArrayBase'

已尝试更换不同JAX版本、切换Colab所有GPU类型,问题仍未解决。

解决方案

1. 解决GPU检测问题

  • 确认Colab已启用GPU:点击菜单栏「修改」→「笔记本设置」,在「硬件加速器」下拉框选择「GPU」,点击「保存」后重启运行时。
  • 验证GPU状态:运行以下代码检查:
!nvidia-smi
import jax
print(jax.devices())

输出包含GpuDevice则说明GPU已正常加载。

2. 修复DeviceArrayBase属性错误

该错误核心原因是新版JAX已移除DeviceArrayBase等旧数组类型,改用统一的jax.Array接口,而当前使用的jaxtorch库代码仍依赖旧版JAX API,有两种解决方式:

方式一:安装兼容的旧版JAX

这是最快捷的方案,直接安装与jaxtorch兼容的JAX版本:

!pip uninstall -y jax jaxlib
!pip install jax==0.3.25 jaxlib==0.3.25+cuda11.cudnn805 -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html

安装完成后重启运行时,再重新执行Notebook代码。

方式二:修改jaxtorch适配代码

若想使用新版JAX,可手动修改库文件:

  1. 打开/content/jax-guided-diffusion/jaxtorch/monkeypatches.py
  2. 将第18行的jaxlib.xla_extension.DeviceArrayBase替换为jax.Array
  3. 将第19行的jax.interpreters.xla.DeviceArray也替换为jax.Array
    保存修改后重新运行导入代码即可。

内容的提问来源于stack exchange,提问作者FrauMoto

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 08:30:26