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

如何让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ó

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 13:37:03