GPU运行PyTorch Lightning代码遇RuntimeError设备不匹配求助
解决PyTorch Lightning中索引张量设备不匹配的RuntimeError
问题场景
在HPC服务器GPU环境下运行PyTorch Lightning代码时触发如下错误:
RuntimeError: indices should be either on cpu or on the same device as the indexed tensor (cpu)
错误触发自代码行:
in_x = self.dataset.data.x[n_id].to(self.device)
原因分析
self.dataset.data.x 存储在CPU设备上,但用于索引的 n_id 已经被转移到了GPU(即self.device指向的设备),PyTorch不允许跨设备进行张量索引操作,因此抛出该错误。
可行解决方案
临时调整索引设备:将索引张量
n_id先移回CPU完成索引,再将结果转到GPU
修改代码为:in_x = self.dataset.data.x[n_id.to('cpu')].to(self.device)提前迁移数据集到GPU(适合小数据集):在模型初始化或setup阶段,把数据集的张量提前移到目标设备
# 示例:在模型的setup方法中添加 def setup(self, stage=None): self.dataset.data.x = self.dataset.data.x.to(self.device)注意:数据集过大时会占用大量GPU显存,需谨慎使用。
规范数据加载流程:检查DataLoader或数据预处理逻辑,确保索引张量
n_id和被索引张量self.dataset.data.x始终处于同一设备,避免提前单独转移索引张量。
内容的提问来源于stack exchange,提问作者Nasim Rony
相关产品推荐
相关产品推荐

