如何让PyTorch Dataset的__getitems__方法返回字典?
PyTorch Dataset中__getitems__返回字典引发KeyError的解决方法
问题代码
import torch import numpy as np from torch.utils.data import Dataset, DataLoader class StackOverflowDataset(torch.utils.data.Dataset): def __init__(self, data): self._data = data def __getitem__(self, idx): return {'item': self._data[idx], 'whatever': idx*self._data[idx]+3} def __getitems__(self, idxs): return {'item': self._data[idxs], 'whatever': idxs*self._data[idxs]+3} def __len__(self): return len(self._data) dataset = StackOverflowDataset(np.random.random(5)) for X in DataLoader(dataset, 2): print(X) break
错误信息
KeyError Traceback (most recent call last) Cell In[182], line 15 12 return len(self._data) 14 dataset = StackOverflowDataset(np.random.random(5)) ---> 15 for X in DataLoader(dataset, 2): 16 print(X) 17 break File ~/recommenders/venv/lib/python3.12/site-packages/torch/utils/data/dataloader.py:630, in _BaseDataLoaderIter.__next__(self) 627 if self._sampler_iter is None: 628 # TODO(https://github.com/pytorch/pytorch/issues/76750) 629 self._reset() # type: ignore[call-arg] --> 630 data = self._next_data() 631 self._num_yielded += 1 632 if self._dataset_kind == _DatasetKind.Iterable and \ 633 self._IterableDataset_len_called is not None and \ 634 self._num_yielded > self._IterableDataset_len_called: File ~/recommenders/venv/lib/python3.12/site-packages/torch/utils/data/dataloader.py:673, in _SingleProcessDataLoaderIter._next_data(self) 671 def _next_data(self): 672 index = self._next_index() # may raise StopIteration --> 673 data = self._dataset_fetcher.fetch(index) # may raise StopIteration 674 if self._pin_memory: 675 data = _utils.pin_memory.pin_memory(data, self._pin_memory_device) File ~/recommenders/venv/lib/python3.12/site-packages/torch/utils/data/_utils/fetch.py:55, in _MapDatasetFetcher.fetch(self, possibly_batched_index) 53 else: 54 data = self.dataset[possibly_batched_index] ---> 55 return self.collate_fn(data) File ~/recommenders/venv/lib/python3.12/site-packages/torch/utils/data/_utils/collate.py:317, in default_collate(batch) 256 def default_collate(batch): 257 r""" 258 Take in a batch of data and put the elements within the batch into a tensor with an additional outer dimension - batch size. 259 (...) 315 >>> default_collate(batch) # Handle `CustomType` automatically 316 """ --> 317 return collate(batch, collate_fn_map=default_collate_fn_map) File ~/recommenders/venv/lib/python3.12/site-packages/torch/utils/data/_utils/collate.py:137, in collate(batch, collate_fn_map) 109 def collate(batch, *, collate_fn_map: Optional[Dict[Union[Type, Tuple[Type, ...]], Callable]] = None): 110 r""" 111 General collate function that handles collection type of element within each batch. 112 (...) 135 for the dictionary of collate functions as `collate_fn_map`. 136 """ --> 137 elem = batch[0] 138 elem_type = type(elem) 140 if collate_fn_map is not None: KeyError: 0
错误原因
你实现的__getitems__直接返回了一个批量级别的字典(每个key对应批量数据),但DataLoader的默认collate_fn期望输入是单个样本的列表(每个样本是__getitem__返回的字典)。当collate_fn拿到这个批量字典时,会把它当成一个单独的样本,尝试通过batch[0]访问第一个元素,而字典中没有key=0,因此抛出KeyError。
解决方法
方法一:返回样本列表(兼容默认collate_fn)
让__getitems__返回与[self.__getitem__(idx) for idx in idxs]等价的结果,即单个样本字典组成的列表:
def __getitems__(self, idxs): return [self.__getitem__(idx) for idx in idxs]
这种方式简单直观,完全兼容默认的collate_fn,不需要修改DataLoader的参数。
方法二:批量构造字典+自定义collate_fn
如果想保留批量处理的效率(避免循环调用__getitem__),可以继续返回批量字典,但需要让collate_fn直接返回该字典,不做额外处理:
# 保留原__getitems__的批量构造逻辑 def __getitems__(self, idxs): return {'item': self._data[idxs], 'whatever': idxs*self._data[idxs]+3} # 创建DataLoader时指定自定义collate_fn for X in DataLoader(dataset, 2, collate_fn=lambda x: x): print(X) break
这种方式利用numpy的批量运算提升效率,自定义的collate_fn直接返回批量字典,避免了默认逻辑的错误处理。
内容的提问来源于stack exchange,提问作者David Davó
相关产品推荐
相关产品推荐

