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

自定义视频数据集返回多帧时DataLoader报NotImplementedError问题

问题描述

我有从YouTube视频中提取的帧袋数据集,希望遍历数据集时返回完整的帧袋。为此编写了自定义VideoDataset类,以及数据集划分、DataLoader构建的相关代码,但在查看DataLoader内容时抛出NotImplementedError,不确定是否允许在__getitem__中返回多图像集合。相关代码与报错信息如下:

自定义数据集类代码

dataset_path = Path('/content/VideoClassificationDataset')

class VideoDataset(Dataset):
  def __init__(self, dictionary, transform = None):
    self.l_dict = list(dictionary.items())
    self.transform = transform
  
  def __len__(self):
    return len(self.l_dict)

  def __get_item__(self, index):
    item = self.l_dict[index]

    images_path = item[0]
    images = [Image.open(f'{dataset_path}/{images_path}/{image}') for image in os.listdir(f'{dataset_path}/{images_path}')]
    
    y_labels = torch.tensor(item[1])

    if self.transform:
      for image in images: self.transform(image)
    
    return images, y_labels

数据集划分与DataLoader构建代码

def spit_train(train_data, perc_val_size):
  train_size = len(train_data)
  val_size = int((train_size * perc_val_size) // 100)
  train_size -= val_size

  return random_split(train_data, [int(train_size), int(val_size)])

train_data, val_data = spit_train(VideoDataset(train_dict, transform=train_transform()), 20)
test_data = VideoDataset(dictionary=test_dict, transform=test_transform())


BATCH_SIZE = 16
NUM_WORKERS = os.cpu_count()

def generate_dataloaders(train_data, test_data, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS):

  train_dl = DataLoader(dataset = train_data, 
                                batch_size = BATCH_SIZE,
                                num_workers = NUM_WORKERS,
                                shuffle = True)

  val_dl = DataLoader(dataset = val_data, 
                                batch_size = BATCH_SIZE,
                                num_workers = NUM_WORKERS,
                                shuffle = True)

  test_dl = DataLoader(dataset = test_data, 
                              batch_size = BATCH_SIZE, 
                              num_workers = NUM_WORKERS, 
                              shuffle = False) # don't need to shuffle testing data when we are considering time series dataset

  return train_dl, val_dl, test_dl

train_dl, val_dl, test_dl = generate_dataloaders(train_data, test_data)

数据集字典示例

{'train/iqGq-8vHEJs/bag_of_shots0': [2],
 'train/iqGq-8vHEJs/bag_of_shots1': [2],
 'train/gnw83R8R6jU/bag_of_shots0': [119],
 'train/gnw83R8R6jU/bag_of_shots1': [119],
...
}

测试代码与报错信息

测试代码:

train_features_batch, train_labels_batch = next(iter(train_dl))
print(train_features_batch.shape, train_labels_batch.shape)

val_features_batch, val_labels_batch = next(iter(val_dl))
print(val_features_batch.shape, val_labels_batch.shape)

报错内容:

NotImplementedError: Caught NotImplementedError in DataLoader worker process 0.
Original Traceback (most recent call last):
  File "/usr/local/lib/python3.8/dist-packages/torch/utils/data/_utils/worker.py", line 302, in _worker_loop
    data = fetcher.fetch(index)
  File "/usr/local/lib/python3.8/dist-packages/torch/utils/data/_utils/fetch.py", line 58, in fetch
    data = [self.dataset[idx] for idx in possibly_batched_index]
  File "/usr/local/lib/python3.8/dist-packages/torch/utils/data/_utils/fetch.py", line 58, in <listcomp>
    data = [self.dataset[idx] for idx in possibly_batched_index]
  File "/usr/local/lib/python3.8/dist-packages/torch/utils/data/dataset.py", line 295, in __getitem__
    return self.dataset[self.indices[idx]]
  File "/usr/local/lib/python3.8/dist-packages/torch/utils/data/dataset.py", line 53, in __getitem__
    raise NotImplementedError
NotImplementedError
解决方案

1. 修复核心拼写错误

报错的直接原因是自定义数据集类里的方法名写错了:__get_item__应该是__getitem__(正确的魔法方法名是双下划线+getitem+双下划线)。PyTorch的Dataset类要求必须实现__getitem__方法,否则会抛出NotImplementedError。

2. 修复transform应用问题

原代码中for image in images: self.transform(image)无效,因为大部分transform不会原地修改图像,而是返回新的变换后图像,应改为:

if self.transform:
    images = [self.transform(image) for image in images]

3. 图像排序处理

os.listdir返回的文件名顺序不确定,会导致同一帧袋的图像加载顺序混乱,需对文件名排序:

image_files = sorted(os.listdir(f'{dataset_path}/{images_path}'))
images = [Image.open(f'{dataset_path}/{images_path}/{image}') for image in image_files]

4. 标签维度调整

原代码生成的标签是形状为[1]的张量,分类任务中通常需要标量或匹配batch的维度,可改为:

y_labels = torch.tensor(item[1]).squeeze()

5. 处理DataLoader的batch拼接

返回图像列表时,默认collate_fn会拼接成大列表而非张量。若每个帧袋帧数相同,可自定义collate_fn转换为张量:

def custom_collate(batch):
    images = [torch.stack(imgs) for imgs, _ in batch]
    images = torch.stack(images)  # 形状: (batch_size, num_frames, C, H, W)
    labels = torch.tensor([label.item() for _, label in batch])
    return images, labels

# 构建DataLoader时指定该函数
train_dl = DataLoader(dataset=train_data, 
                      batch_size=BATCH_SIZE,
                      num_workers=NUM_WORKERS,
                      shuffle=True,
                      collate_fn=custom_collate)

修复后的完整数据集类

dataset_path = Path('/content/VideoClassificationDataset')

class VideoDataset(Dataset):
    def __init__(self, dictionary, transform=None):
        self.l_dict = list(dictionary.items())
        self.transform = transform
    
    def __len__(self):
        return len(self.l_dict)

    def __getitem__(self, index):
        item = self.l_dict[index]
        images_path = item[0]
        
        # 排序文件名保证加载顺序一致
        image_files = sorted(os.listdir(f'{dataset_path}/{images_path}'))
        images = [Image.open(f'{dataset_path}/{images_path}/{image}') for image in image_files]
        
        y_labels = torch.tensor(item[1]).squeeze()

        if self.transform:
            images = [self.transform(image) for image in images]
        
        return images, y_labels

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 09:02:51