TPU V4-16/64初始化失败:无法建立SliceBuilder gRPC通道求助
TPU V4-16/V4-64初始化失败,无法运行训练代码
问题背景
拥有32个按需免费TPU V4芯片,尝试用TPU V4-64训练模型,但V4-8可正常运行,V4-16至V4-64时,官方教程代码和自研代码均无法运行。
操作步骤
严格遵循GCP TPU官方教程配置torch_xla,步骤如下:
- 创建TPU VM:
gcloud compute tpus queued-resources create RESOURCE_NAME \ --node-id NODE_NAME \ --project PROJECT_NAME \ --zone us-central2-b \ --accelerator-type v4-64 \ --runtime-version tpu-ubuntu2204-base
- 连接TPU VM:
gcloud compute tpus tpu-vm ssh NODE_NAME --zone=us-central2-b
- 安装依赖:
pip install torch~=2.3.0 torch_xla[tpu]~=2.3.0 torchvision -f https://storage.googleapis.com/libtpu-releases/index.html
- 克隆torch_xla仓库:
git clone --depth=1 --branch r2.3 https://github.com/pytorch/xla.git
- 运行测试代码:
PJRT_DEVICE=TPU python3 xla/test/test_train_mp_imagenet.py --fake_data --batch_size=256 --num_epochs=1
尝试结果
尝试1:运行测试代码
长时间无进度输出,V4-8可正常利用xla::0~xla::3并打印训练/测试损失,日志片段:
WARNING:root:PJRT is now the default runtime. For more information, see https://github.com/pytorch/xla/blob/master/docs/pjrt.md WARNING:root:libtpu.so and TPU device found. Setting PJRT_DEVICE=TPU. WARNING: All log messages before absl::InitializeLog() is called are written to STDERR I0000 00:00:1720161247.044048 21553 pjrt_api.cc:100] GetPjrtApi was found for tpu at /home/hungwon3626/.local/lib/python3.10/site-packages/libtpu/libtpu.so I0000 00:00:1720161247.044129 21553 pjrt_api.cc:79] PJRT_Api is set for device type tpu I0000 00:00:1720161247.044136 21553 pjrt_api.cc:146] The PJRT plugin has PJRT API version 0.46. The framework PJRT API version is 0.46. ...
尝试2:手动初始化TPU设备
暂停运行后执行以下代码:
>>> import os >>> os.environ['TPU_NUM_DEVICES'] = '32' # 未设置该环境变量重试后结果一致 >>> import torch_xla as xla >>> xla.device()
出现关键错误:
WARNING:root:PJRT is now the default runtime. For more information, see https://github.com/pytorch/xla/blob/master/docs/pjrt.md WARNING:root:libtpu.so and TPU device found. Setting PJRT_DEVICE=TPU. WARNING: All log messages before absl::InitializeLog() is called are written to STDERR I0000 00:00:1720161247.044048 21553 pjrt_api.cc:100] GetPjrtApi was found for tpu at /home/hungwon3626/.local/lib/python3.10/site-packages/libtpu/libtpu.so I0000 00:00:1720161247.044129 21553 pjrt_api.cc:79] PJRT_Api is set for device type tpu I0000 00:00:1720161247.044136 21553 pjrt_api.cc:146] The PJRT plugin has PJRT API version 0.46. The framework PJRT API version is 0.46. Traceback (most recent call last): File "<stdin>", line 1, in <module> File "/home/hungwon3626/.local/lib/python3.10/site-packages/torch_xla/torch_xla.py", line 21, in device return xm.xla_device(index) File "/home/hungwon3626/.local/lib/python3.10/site-packages/torch_xla/core/xla_model.py", line 212, in xla_device return runtime.xla_device(n, devkind) File "/home/hungwon3626/.local/lib/python3.10/site-packages/torch_xla/runtime.py", line 95, in wrapper return fn(*args, **kwargs) File "/home/hungwon3626/.local/lib/python3.10/site-packages/torch_xla/runtime.py", line 124, in xla_device return torch.device(torch_xla._XLAC._xla_get_default_device()) RuntimeError: Bad StatusOr access: UNKNOWN: TPU initialization failed: Failed to establish SliceBuilder grpc channel to 10.130.15.203:8471.
期望目标
利用全部32个芯片在XLA设备上训练单个大型模型。
依赖版本
torch 2.3.1 torch-xla 2.3.0 torchvision 0.18.1 Python 3.10.6
解决方案
1. 排查TPU节点网络连通性
错误提示Failed to establish SliceBuilder grpc channel说明VM与TPU切片管理器的网络连接异常:
- 在TPU VM内执行
ping 10.130.15.203,检查是否能连通目标IP - 执行
telnet 10.130.15.203 8471,确认8471端口是否开放
如果连通失败,删除当前TPU节点重新创建,可尝试切换至us-central2-a等其他可用区,部分区域可能存在临时网络故障。
2. 显式配置多进程环境变量
对于V4-16及以上多芯片TPU,需显式设置多进程相关环境变量,运行训练代码前执行:
export PJRT_DEVICE=TPU export XLA_USE_BF16=1 export NUM_TPU_WORKERS=32 export TPU_NUM_DEVICES=32
再重新运行测试代码:
python3 xla/test/test_train_mp_imagenet.py --fake_data --batch_size=256 --num_epochs=1
3. 升级libtpu版本
当前torch-xla 2.3.0对应的libtpu可能存在V4大切片兼容问题,执行以下命令升级:
pip install --upgrade libtpu-nightly
升级后重启Python进程再尝试初始化TPU。
4. 使用官方预配置镜像
避免手动安装的版本冲突,创建TPU VM时使用预配置好兼容依赖的镜像:
gcloud compute tpus queued-resources create RESOURCE_NAME \ --node-id NODE_NAME \ --project PROJECT_NAME \ --zone us-central2-b \ --accelerator-type v4-64 \ --runtime-version tpu-vm-pytorch-2.3
5. 验证TPU设备可用性
在VM内执行以下代码,确认能检测到所有TPU设备:
import torch_xla.core.xla_model as xm devices = xm.get_xla_supported_devices() print(f"Detected {len(devices)} TPU devices: {devices}")
如果能输出32个设备,说明初始化成功,再运行训练代码。
内容的提问来源于stack exchange,提问作者gnsrnjs
相关产品推荐
相关产品推荐

