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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 21:17:38