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

在Google Colab TPU上运行PyTorch 2.2遇报错,求解决方案

在Google Colab TPU上运行PyTorch 2.2时的JAX警告及TPU初始化失败问题

问题概述

在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.

解决步骤

  1. 确认Colab TPU配置
    打开Colab菜单栏「修改」→「笔记本设置」,确认硬件加速器选择「TPU」,保存后重启运行时环境。

  2. 重新安装适配的依赖包
    替换原安装命令为以下命令,强制重新安装适配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
    
  3. 提前初始化TPU环境变量
    在导入torch_xla前添加环境变量配置代码:

    import os
    os.environ['XLA_USE_BF16'] = '1'
    os.environ['TPU_NAME'] = 'grpc://' + os.environ['COLAB_TPU_ADDR']
    
  4. 正确获取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)
    
  5. 处理JAX警告
    torch_xla底层依赖JAX实现TPU通信,因此未主动使用JAX也会出现警告。若要消除警告,可升级JAX版本:

    !pip install --upgrade jax jaxlib
    

    若无需消除,只要TPU正常初始化,该警告不影响使用。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 21:45:35