自定义图数据集迭代批量获取报错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
相关产品推荐
相关产品推荐

