PyTorch DataLoader迭代器next方法失效问题排查
问题
我构建了一个DataLoader并加载了一个仅含10个样本(标签为0至9)的简单数据集来测试迭代器,但调用next方法时迭代器无法正常推进,请问代码存在什么问题?
import torch from torch.utils.data import Dataset from torchvision import datasets from torchvision.transforms import ToTensor from torch.utils.data import DataLoader class CustomImageDataset(Dataset): def __init__(self, features, labels, transform=None, target_transform=None): self.labels = labels self.features = features def __len__(self): return len(self.labels) def __getitem__(self, idx): image = self.features[idx] label = self.labels[idx] return image, label features = [0,1,2,3,4,5,6,7,8,9] labels = [0,1,2,3,4,5,6,7,8,9] dataset = CustomImageDataset(features, labels) dataloader = DataLoader(dataset, batch_size=2, shuffle=False, num_workers=0) for i in range(20): features, labels = next(iter(dataloader)) print("=========") print('This is ' + str(i + 1) + "th time.") print(labels[0]) print(labels[1])
运行结果:
========= This is 1th time. tensor(0) tensor(1) ========= This is 2th time. tensor(0) tensor(1) ========= This is 3th time. tensor(0) tensor(1) ========= This is 4th time. tensor(0) tensor(1) ========= This is 5th time. tensor(0) tensor(1) ========= This is 6th time. tensor(0) tensor(1)
问题分析与解决
问题核心在于循环内的next(iter(dataloader)):每次调用iter(dataloader)都会生成一个全新的迭代器,每次next取的都是这个新迭代器的第一个batch,所以永远输出[0,1],不会推进迭代。
两种修正方案:
方案一:提前创建一次迭代器,循环调用next
先初始化一次迭代器,之后每次next都基于同一个迭代器推进;当迭代器耗尽时,重新生成迭代器即可继续遍历:
import torch from torch.utils.data import Dataset from torch.utils.data import DataLoader class CustomImageDataset(Dataset): def __init__(self, features, labels, transform=None, target_transform=None): self.labels = labels self.features = features def __len__(self): return len(self.labels) def __getitem__(self, idx): image = self.features[idx] label = self.labels[idx] return image, label features = [0,1,2,3,4,5,6,7,8,9] labels = [0,1,2,3,4,5,6,7,8,9] dataset = CustomImageDataset(features, labels) dataloader = DataLoader(dataset, batch_size=2, shuffle=False, num_workers=0) # 提前创建迭代器 data_iter = iter(dataloader) for i in range(20): try: features, labels = next(data_iter) except StopIteration: # 迭代器耗尽后重新生成 data_iter = iter(dataloader) features, labels = next(data_iter) print("=========") print(f'This is {i + 1}th time.') print(labels[0]) print(labels[1])
方案二:直接遍历DataLoader(推荐)
PyTorch的DataLoader本身是可迭代对象,直接用for循环遍历会自动处理迭代器的创建和耗尽后的重置,代码更简洁:
import torch from torch.utils.data import Dataset from torch.utils.data import DataLoader class CustomImageDataset(Dataset): def __init__(self, features, labels, transform=None, target_transform=None): self.labels = labels self.features = features def __len__(self): return len(self.labels) def __getitem__(self, idx): image = self.features[idx] label = self.labels[idx] return image, label features = [0,1,2,3,4,5,6,7,8,9] labels = [0,1,2,3,4,5,6,7,8,9] dataset = CustomImageDataset(features, labels) dataloader = DataLoader(dataset, batch_size=2, shuffle=False, num_workers=0) count = 0 # 无限遍历直到达到20次 while count < 20: for features, labels in dataloader: if count >=20: break count +=1 print("=========") print(f'This is {count}th time.') print(labels[0]) print(labels[1])
内容的提问来源于stack exchange,提问作者X.G
相关产品推荐
相关产品推荐

