在Google Colab TPU上运行PyTorch 2.2遇报错,求解决方案
问题概述
在Google Colab的TPU环境中运行PyTorch 2.2,未主动使用JAX库却出现JAX相关初始化警告,且尝试使用TPU设备时触发初始化失败错误。
操作流程
安装依赖
执行以下命令安装PyTorch及torch_xla:
!pip install torch~=2.2.0 torch_xla[tpu]~=2.2.0 -f https://storage.googleapis.com/libtpu-releases/index.html
导入模块
运行代码导入PyTorch和torch_xla:
import torch import torch_xla.core.xla_model as xm
报错详情
导入时的警告信息
/usr/local/lib/python3.10/dist-packages/jax/init.py:27: UserWarning: cloud_tpu_init failed: KeyError('')
This a JAX bug; please report an issue at https://github.com/google/jax/issues
_warn(f"cloud_tpu_init failed: {repr(exc)}\n This a JAX bug; please report "
/usr/local/lib/python3.10/dist-packages/transformers/utils/generic.py:441: UserWarning: torch.utils._pytree._register_pytree_node is deprecated. Please use torch.utils._pytree.register_pytree_node instead.
_torch_pytree._register_pytree_node(
TPU初始化失败错误
运行以下张量运算代码:
t1 = torch.tensor(100, device=xm.xla_device()) t2 = torch.tensor(200, device=xm.xla_device()) print(t1 + t2)
触发错误:
2 frames /usr/local/lib/python3.10/dist-packages/torch_xla/runtime.py in xla_device(n, devkind) 121 122 if n is None: --> 123 return torch.device(torch_xla._XLAC._xla_get_default_device()) 124 125 devices = xm.get_xla_supported_devices(devkind=devkind) RuntimeError: Bad StatusOr access: UNKNOWN: TPU initialization failed: No ba16c7433 device found.
解决步骤
确认Colab TPU配置
打开Colab菜单栏「修改」→「笔记本设置」,确认硬件加速器选择「TPU」,保存后重启运行时环境。重新安装适配的依赖包
替换原安装命令为以下命令,强制重新安装适配Colab TPU的版本:!pip install torch==2.2.0 torch_xla[tpu]~=2.2.0 -f https://storage.googleapis.com/libtpu-releases/index.html --force-reinstall提前初始化TPU环境变量
在导入torch_xla前添加环境变量配置代码:import os os.environ['XLA_USE_BF16'] = '1' os.environ['TPU_NAME'] = 'grpc://' + os.environ['COLAB_TPU_ADDR']正确获取TPU设备
替换原设备获取代码,改用get_xla_supported_devices()获取设备列表:devices = xm.get_xla_supported_devices() t1 = torch.tensor(100, device=devices[0]) t2 = torch.tensor(200, device=devices[0]) print(t1 + t2)处理JAX警告
torch_xla底层依赖JAX实现TPU通信,因此未主动使用JAX也会出现警告。若要消除警告,可升级JAX版本:!pip install --upgrade jax jaxlib若无需消除,只要TPU正常初始化,该警告不影响使用。
内容的提问来源于stack exchange,提问作者Thomas Johnson

