使用PyG处理数据集时遇ImportError:无法导入dropout_edge
问题描述
运行PyG数据集预处理代码时出现错误:
dataset_sample = OneStepDataset(OUTPUT_DIR, "valid", return_pos=True) graph, position = dataset_sample[0]
其中OneStepDataset是torch_geometric.data.Dataset的子类,报错追踪到PyG源码dataset.py第197行:
if (isinstance(idx, (int, np.integer)) 194 or (isinstance(idx, Tensor) and idx.dim() == 0) 195 or (isinstance(idx, np.ndarray) and np.isscalar(idx))): --> 197 data = self.get(self.indices()[idx]) 198 data = data if self.transform is None else self.transform(data) 199 return data
最终报错信息:
ImportError: cannot import name 'dropout_edge' from 'torch_geometric.utils' (c:\Users\...\AppData\Local\Programs\Python\Python39\lib\site-packages\torch_geometric\utils\__init__.py)
解决方案
- 检查PyG版本:执行命令
pip show torch_geometric,查看Version字段。dropout_edge是PyG 0.8.0及以上版本新增的API,旧版本无此函数。 - 升级PyG兼容版本:执行
pip install --upgrade torch_geometric,需保证PyTorch版本与PyG版本匹配。若出现依赖错误,单独安装相关依赖包,命令格式如下(替换${TORCH}为你的PyTorch版本,${CUDA}为你的CUDA版本,如cpu、cu113):
pip install torch_scatter torch_sparse torch_cluster torch_spline_conv -f https://data.pyg.org/whl/torch-${TORCH}+${CUDA}.html
- 无法升级时的替代方案:自行实现
dropout_edge功能,示例代码如下:
import torch def dropout_edge(edge_index, p=0.5, training=True): if not training or p == 0.0: return edge_index, torch.ones(edge_index.size(1), dtype=torch.bool) mask = torch.rand(edge_index.size(1), device=edge_index.device) >= p return edge_index[:, mask], mask
内容的提问来源于stack exchange,提问作者playerJX1
相关产品推荐
相关产品推荐

