继承torch_geometric Dataset时无法实例化抽象类的问题
PyTorch Geometric Dataset类实例化错误问题解答
问题代码
from torch_geometric.data import Data, Dataset import numpy as np class GraphTimeSeriesDataset(Dataset): def __init__(self): # Initialize your dataset here pass def __len__(self): # Return the total number of samples in the dataset return len(self.data_list) def __getitem__(self, idx): # Return the data sample at the given index return self.data_list[idx] dataset = GraphTimeSeriesDataset()
错误信息
--------------------------------------------------------------------------- TypeError Traceback (most recent call last) Cell In[4], line 17 13 def __getitem__(self, idx): 14 # Return the data sample at the given index 15 return self.data_list[idx] ---> 17 dataset = GraphTimeSeriesDataset() TypeError: Can't instantiate abstract class GraphTimeSeriesDataset with abstract methods get, len
问题解答
1. 错误由环境还是代码本身导致?
由环境版本差异导致,代码在特定版本的PyTorch Geometric中可正常运行,但版本不同对类的实现要求不一致。
2. 如何调整代码避免错误?
有两种适配方案:
- 方案一:适配新版本要求,重写
len()和get()方法(替代原魔术方法)from torch_geometric.data import Data, Dataset import numpy as np class GraphTimeSeriesDataset(Dataset): def __init__(self): super().__init__() # 必须调用父类初始化方法 # 此处替换为实际数据加载逻辑,示例初始化空列表 self.data_list = [] def len(self): return len(self.data_list) def get(self, idx): return self.data_list[idx] dataset = GraphTimeSeriesDataset() - 方案二:兼容新旧版本,继承
InMemoryDataset(更适合内存型数据集)from torch_geometric.data import InMemoryDataset, Data import numpy as np class GraphTimeSeriesDataset(InMemoryDataset): def __init__(self): super().__init__('./dataset_cache') # 指定数据集缓存路径 # 加载并处理数据为Data对象列表 self.data_list = [] self.data, self.slices = self.collate(self.data_list) def process(self): # 在此实现数据读取、转换为Data对象的逻辑 pass
3. 根本原因是什么?
PyTorch Geometric的Dataset抽象类在版本迭代中修改了强制实现的方法:
- 旧版本接受子类实现
__len__和__getitem__魔术方法; - 新版本将强制实现的抽象方法改为
len()和get(),未实现这两个方法的子类会被判定为抽象类,无法实例化。
另外原代码未调用父类__init__方法,在旧版本可能未触发报错,但新版本中这是必要操作。
内容的提问来源于stack exchange,提问作者theabc50111
相关产品推荐
相关产品推荐

