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
相关产品推荐
相关产品推荐

