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

Colab TPU运行PyTorch Lightning报错:输入张量非XLA张量

问题原因

你用torch.arange()手动创建的pos张量默认会在CPU上生成——PyTorch的张量创建函数(比如torch.arange、torch.tensor)如果不指定device参数,都会默认使用CPU设备。而你的模型参数(如tok_embedding、pos_embedding)已经被PyTorch Lightning自动部署到TPU,输入src也在TPU上,此时用CPU的pos调用TPU上的嵌入层,就会触发"Input Tensor is not an XLA tensor"错误(TPU使用XLA类型张量,CPU张量不属于该类型)。

另外注意你的self.scale也是用torch.FloatTensor在CPU创建的,后续和TPU张量计算时也会出问题,需要一并处理。

解决办法
  • 对齐pos和输入src的设备
    创建pos后直接转移到src所在设备,适配任意设备场景:

    pos = torch.arange(0, src_len).unsqueeze(0).repeat(batch_size, 1).to(src.device)
    
  • 利用LightningModule的self.device指定设备
    因为Encoder继承了pl.LightningModule,可以直接用self.device绑定模型所在设备:

    # 创建时直接指定设备
    pos = torch.arange(0, src_len, device=self.device).unsqueeze(0).repeat(batch_size, 1)
    # 或者创建后转移
    pos = torch.arange(0, src_len).unsqueeze(0).repeat(batch_size, 1).to(self.device)
    
  • 修复self.scale的设备问题
    在__init__中将self.scale转移到模型设备:

    self.scale = torch.sqrt(torch.FloatTensor([hid_dim])).to(self.device)
    

    或者创建时直接指定设备:

    self.scale = torch.sqrt(torch.tensor([hid_dim], device=self.device, dtype=torch.float32))
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 01:20:36