PyTorch Geometric调用RandomLinkSplit报张量设备不一致错误
问题原因
该问题是PyTorch Geometric v2.0.2版本的已知缺陷:RandomLinkSplit内部调用负采样逻辑生成负样本时,没有自动匹配输入数据所在的CUDA设备,默认会将生成的负样本张量放在CPU上。你传入的原始数据所有张量都在cuda:0,但内部生成的负样本在CPU,正负样本拼接时就触发了设备不一致的报错。
解决方案
- 方案1(最推荐,无需修改依赖):调整数据转移到GPU的顺序,先在CPU上完成链接拆分,再将拆分后的数据集转移到CUDA设备。拆分操作本身计算量极低,CPU运行不会有性能损失,修改后的代码示例如下:
import torch import torch_geometric.transforms as T from torch_geometric.nn import GCNConv from torch_geometric.datasets import Planetoid from torch_geometric.utils import negative_sampling device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # 初始转换不包含设备转移 transform = T.Compose([ T.NormalizeFeatures(), ]) dataset = Planetoid(root='/tmp/Planetoid', name='Cora', transform=transform) data = dataset[0] # 先在CPU上完成链接拆分 split_transform = T.RandomLinkSplit(num_val=0.05, num_test=0.1, is_undirected=True,) train_data, val_data, test_data = split_transform(data) # 拆分完成后分别转移到目标设备 train_data = train_data.to(device) val_data = val_data.to(device) test_data = test_data.to(device)
- 方案2:升级PyTorch Geometric到v2.1.0及以上版本,官方已在后续版本修复了该设备不匹配问题,升级后原有代码可直接正常运行。
- 方案3:如果必须保留当前版本,可手动修改
RandomLinkSplit源码,找到内部调用negative_sampling的代码行,新增参数device=edge_index.device,强制负样本生成在和输入边相同的设备上,该方式不推荐。
内容的提问来源于stack exchange,提问作者Sticky
相关产品推荐
相关产品推荐

