You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

继承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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.21 01:02:21