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

PyTorch自定义数据集报错:list indices must be integers or slices, not list

PyTorch自定义Dataset调用异常:外部文件类接收列表idx,Notebook内类正常

问题场景

  • 模型训练10、100轮epoch均正常,调整为500轮后GPU崩溃
  • 重启GPU后Jupyter Notebook抛出500服务器错误,执行pip install --upgrade nbconvert修复后,原代码无法运行
  • 调试差异:
    • 自定义数据集类放在外部Python文件(src/custom_dataset.py)中调用时,__getitem__的参数idx为列表,触发TypeError
    • 同一类直接定义在Notebook单元格内调用时,idx为整数,代码运行正常

相关代码

外部文件src/custom_dataset.py中的数据集类

import os
from natsort import natsorted
from PIL import Image
from datasets import Dataset

class LoadPairedDataset(Dataset):
    def __init__(self, root_dir, transform=None):
        self.root_dir = root_dir
        self.transform = transform
        self.images = os.listdir(root_dir)

    def __len__(self):
        return len(self.images)

    def __getitem__(self, idx):
        print(idx)
        img_name = os.path.join(self.root_dir, self.images[idx])
        image = Image.open(img_name)

        if self.transform:
            image = self.transform(image)

        return image

Jupyter Notebook调用代码

# Imports
import os
from PIL import Image
import torch
from torch.utils.data import Dataset
from torchvision import transforms
from src.custom_dataset import LoadPairedDataset
 
# Notebook内定义的数据集类
class CustomImageDataset(Dataset):
    def __init__(self, root_dir, transform=None):
        self.root_dir = root_dir
        self.transform = transform
        self.images = os.listdir(root_dir)

    def __len__(self):
        return len(self.images)

    def __getitem__(self, idx):
        print(idx)  
        img_name = os.path.join(self.root_dir, self.images[idx])
        image = Image.open(img_name)

        if self.transform:
            image = self.transform(image)

        return image
    

base_path = "../lol-custom"

transform = transforms.Compose([
            transforms.ToTensor()                       
        ])

# Notebook内类调用(正常运行)
train_data = CustomImageDataset(root_dir=base_path + "/train/low", transform=transform)
dataloader = torch.utils.data.DataLoader(
    train_data,
    batch_size=5,
    sampler=None,
    num_workers=0
)
print(next(iter(dataloader)).shape) # 输出0 1 2 3 4,以及torch.Size([5, 3, 400, 600])
print("######### 以下调用会报错 ##############")

# 外部文件类调用(报错)
train_data = LoadPairedDataset(root_dir=base_path + "/train/low", transform=transform)
dataloader = torch.utils.data.DataLoader(
    train_data,
    batch_size=5,
    sampler=None,
    num_workers=0
)
print(next(iter(dataloader)).shape) 

环境信息

  • PyTorch 2.0.1
  • Python 3.9.18

报错信息

0
1
2
3
4
torch.Size([5, 3, 400, 600])
######### 以下调用会报错 ##############
[0, 1, 2, 3, 4]

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
Cell In[1], line 61
     52 dataloader = torch.utils.data.DataLoader(
     53     train_data,
     54     batch_size=5,
     55     sampler=None,
     56     num_workers=0
     57 )
     58 # 输出[0, 1, 2, 3, 4]后触发错误
---> 61 print(next(iter(dataloader)).shape)

File ~\anaconda3\envs\mmie\lib\site-packages\torch\utils\data\dataloader.py:633, in _BaseDataLoaderIter.__next__(self)
    630 if self._sampler_iter is None:
    631     # TODO(https://github.com/pytorch/pytorch/issues/76750)
    632     self._reset()  # type: ignore[call-arg]
--> 633 data = self._next_data()
    634 self._num_yielded += 1
    635 if self._dataset_kind == _DatasetKind.Iterable and \
    636         self._IterableDataset_len_called is not None and \
    637         self._num_yielded > self._IterableDataset_len_called:

File ~\anaconda3\envs\mmie\lib\site-packages\torch\utils\data\dataloader.py:677, in _SingleProcessDataLoaderIter._next_data(self)
    675 def _next_data(self):
    676     index = self._next_index()  # may raise StopIteration
--> 677     data = self._dataset_fetcher.fetch(index)  # may raise StopIteration
    678     if self._pin_memory:
    679         data = _utils.pin_memory.pin_memory(data, self._pin_memory_device)

File ~\anaconda3\envs\mmie\lib\site-packages\torch\utils\data\_utils\fetch.py:49, in _MapDatasetFetcher.fetch(self, possibly_batched_index)
     47 if self.auto_collation:
     48     if hasattr(self.dataset, "__getitems__") and self.dataset.__getitems__:
---> 49         data = self.dataset.__getitems__(possibly_batched_index)
     50     else:
     51         data = [self.dataset[idx] for idx in possibly_batched_index]

File ~\anaconda3\envs\mmie\lib\site-packages\datasets\arrow_dataset.py:2807, in Dataset.__getitems__(self, keys)
   2805 def __getitems__(self, keys: List) -> List:
   2806     """Can be used to get a batch using a list of integers indices."""
-> 2807     batch = self.__getitem__(keys)
   2808     n_examples = len(batch[next(iter(batch))])
   2809     return [{col: array[i] for col, array in batch.items()} for i in range(n_examples)]

File ~\Projects\mmie\src\custom_dataset.py:21, in LoadPairedDataset.__getitem__(self, idx)
     19 def __getitem__(self, idx):
     20     print(idx)
---> 21     img_name = os.path.join(self.root_dir, self.images[idx])
     22     image = Image.open(img_name)
     24     if self.transform:

TypeError: list indices must be integers or slices, not list

问题根源与解决方案

根源

外部文件中的LoadPairedDataset错误继承了Hugging Face datasets库的Dataset类,而非PyTorch原生的torch.utils.data.Dataset。

PyTorch的DataLoader会自动检测数据集是否实现了__getitems__方法:

  • Hugging Face的Dataset类内置了__getitems__,DataLoader会直接传入批量索引列表(如[0,1,2,3,4])
  • 未实现__getitems__的PyTorch原生Dataset子类,DataLoader会循环调用__getitem__并传入单个整数索引

解决方案

修改src/custom_dataset.py中的导入和继承,替换为PyTorch原生的Dataset:

import os
from natsort import natsorted
from PIL import Image
# 替换导入为PyTorch的Dataset
from torch.utils.data import Dataset

class LoadPairedDataset(Dataset):
    def __init__(self, root_dir, transform=None):
        self.root_dir = root_dir
        self.transform = transform
        self.images = os.listdir(root_dir)

    def __len__(self):
        return len(self.images)

    def __getitem__(self, idx):
        print(idx)
        img_name = os.path.join(self.root_dir, self.images[idx])
        image = Image.open(img_name)

        if self.transform:
            image = self.transform(image)

        return image

修改后重新运行代码,__getitem__将接收整数索引,与Notebook内的类行为一致,不再触发TypeError。

内容的提问来源于stack exchange,提问作者stic-lab

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 19:50:54