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

自定义图数据集迭代批量获取报错KeyError:0的原因与修复方案

问题描述

自定义数据集仅包含特征矩阵和邻接矩阵:特征矩阵是4000×3000的张量,edge_index尺寸为torch.Size([2, 18708]),代码如下:

from torch_geometric.data import Data
data = Data(x=features,edge_index=edge_index.t().contiguous())
from torch_geometric.data import DataLoader
loader = DataLoader(data,batch_size=128,shuffle=True)
for batch in loader:
    batch

运行后出现报错:

File ~/miniconda3/envs/spacel/lib/python3.8/site-packages/torch/utils/data/dataloader.py:521, in _BaseDataLoaderIter.__next__(self)
    519 if self._sampler_iter is None:
    520     self._reset()
--> 521 data = self._next_data()
    522 self._num_yielded += 1
    523 if self._dataset_kind == _DatasetKind.Iterable and \
    524         self._IterableDataset_len_called is not None and \
    525         self._num_yielded > self._IterableDataset_len_called:

File ~/miniconda3/envs/spacel/lib/python3.8/site-packages/torch/utils/data/dataloader.py:561, in _SingleProcessDataLoaderIter._next_data(self)
    559 def _next_data(self):
    560     index = self._next_index()  # may raise StopIteration
--> 561 data = self._dataset_fetcher.fetch(index)  # may raise StopIteration
    562 if self._pin_memory:
    563         data = _utils.pin_memory.pin_memory(data)

File ~/miniconda3/envs/spacel/lib/python3.8/site-packages/torch/utils/data/_utils/fetch.py:49, in _MapDatasetFetcher.fetch(self, possibly_batched_index)
     47 def fetch(self, possibly_batched_index):
     48     if self.auto_collation:
...
File ~/miniconda3/envs/spacel/lib/python3.8/site-packages/torch_geometric/data/storage.py:111, in BaseStorage.__getitem__(self, key)
    110 def __getitem__(self, key: str) -> Any:
--> 111     return self._mapping[key]

KeyError: 0
报错原因
  • DataLoader的输入必须是数据集对象(如Dataset子类、包含多个Data实例的列表),但你直接传入了单个Data对象。
  • 单个Data实例不支持按索引(如data[0])访问,DataLoader尝试用索引0获取数据时触发了KeyError。
修复方法

根据你的任务场景选择对应的修复方式:

场景1:单图节点级任务(按节点批量加载)

如果是节点分类/回归这类单图上的节点级任务,普通DataLoader不适用,需用PyG专门的节点采样加载器,比如NeighborLoader:

from torch_geometric.data import Data
from torch_geometric.loader import NeighborLoader

data = Data(x=features, edge_index=edge_index.t().contiguous())
# 配置节点采样加载器,batch_size为每次采样的节点数,num_neighbors为每层采样的邻居数量
loader = NeighborLoader(
    data,
    batch_size=128,
    shuffle=True,
    num_neighbors=[20, 10],  # 可根据任务需求调整邻居采样数
    input_nodes=None,  # None表示对所有节点进行采样
)

for batch in loader:
    print(batch)

场景2:多图任务(数据集包含多个图)

如果你的数据集实际由多个图组成,只是误将所有图的特征和边合并成了单个Data,需拆分出每个图的Data实例,放入列表后再用DataLoader:

from torch_geometric.data import Data, DataLoader

# 假设已拆分得到多个图的Data实例,存入data_list
data_list = [
    Data(x=graph1_x, edge_index=graph1_edge_index),
    Data(x=graph2_x, edge_index=graph2_edge_index),
    # ... 其他图的Data实例
]
loader = DataLoader(data_list, batch_size=128, shuffle=True)

for batch in loader:
    print(batch)

场景3:仅测试单图加载(无节点采样)

如果只是想将整个图作为一个batch加载,只需将单个Data实例包装成列表:

from torch_geometric.data import Data, DataLoader

data = Data(x=features, edge_index=edge_index.t().contiguous())
loader = DataLoader([data], batch_size=1, shuffle=True)

for batch in loader:
    print(batch)

内容的提问来源于stack exchange,提问作者luxin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 10:47:04