自定义视频数据集返回多帧时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
相关产品推荐
相关产品推荐

