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

在Google Colab中无法获取带requires_grad的TPU上PyTorch张量值

解决Colab中TPU上带requires_grad的PyTorch张量无法取值的问题

问题复现

在Google Colab使用TPU运行PyTorch时,单个带requires_grad=True的TPU张量可以正常打印,但对多个这类张量进行运算后,尝试打印结果、调用t.item()或t.cpu().item()时,程序会无限冻结:

单个张量正常运行代码:

import torch
import torch_xla.core.xla_model as xm
device = xm.xla_device()

t = torch.tensor([1.], requires_grad=True).to(device)
print(t)

运算后冻结的代码:

import torch
import torch_xla.core.xla_model as xm
device = xm.xla_device()

t1 = torch.tensor([1.], requires_grad=True).to(device)
t2 = torch.tensor([1.], requires_grad=True).to(device)
t3 = t1 + t2
print(t3)

问题原因

PyTorch XLA采用**延迟执行(Lazy Execution)**机制,所有TPU上的运算会先构建计算图,直到显式触发同步操作才会实际执行。当张量开启梯度追踪时,运算后的计算图未被同步,导致取值操作无法获取计算结果,进而卡住程序。

解决方案

在需要获取张量值之前,调用xm.mark_step()触发计算图的执行和同步:

修正后的代码示例

import torch
import torch_xla.core.xla_model as xm
device = xm.xla_device()

t1 = torch.tensor([1.], requires_grad=True).to(device)
t2 = torch.tensor([1.], requires_grad=True).to(device)
t3 = t1 + t2

# 触发计算图执行与同步
xm.mark_step()

# 现在可以正常获取值
print(t3)
print(t3.item())
print(t3.cpu().item())

补充说明

  • xm.mark_step()会强制TPU执行当前所有待处理的运算,并同步结果到主机端,确保后续的取值操作能拿到实际计算结果。
  • 如果后续还需要对张量进行反向传播,xm.mark_step()不会破坏计算图,依然可以正常调用t3.backward()等梯度操作。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 20:59:54