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

为何自定义Dataset的PyTorch DataLoader输出结构与官方示例不符?

DataLoader输出结构与官方DCGAN教程不符的问题

我在实现PyTorch官方DCGAN人脸教程时,遇到了DataLoader输出结构不一致的问题:

官方示例中的DataLoader调用与结构

官方代码里,调用DataLoader的方式如下:

real_batch = next(iter(dataloader))
real_batch[0].to(device)[:64]

或者循环中:

for i, data in enumerate(dataloader, 0):
    real_cpu = data[0].to(device)

这里的data[0]对应大小为预设batch_size(比如128)的整批样本,示例中用[:64]截取前64个。

官方的Dataset与DataLoader定义:

dataset = dset.ImageFolder(root=dataroot,
                           transform=transforms.Compose([
                               transforms.Resize(image_size),
                               transforms.CenterCrop(image_size),
                               transforms.ToTensor(),
                               transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
                           ]))
dataloader = torch.utils.data.DataLoader(dataset, batch_size=batch_size,
                                         shuffle=True, num_workers=workers)

我的自定义Dataset与问题

我自己定义的Dataset及DataLoader如下:

class MyDataset(torch.utils.data.Dataset):
    def __init__(self, dataset, image_size):
        super(MyDataset, self).__init__()
        # dataset是单通道PIL图像组成的列表
        self.dataset = dataset
        self.image_size = image_size
        self.transform=transforms.Compose([transforms.Resize(self.image_size),
                               transforms.ToTensor(),
                               transforms.Normalize((0.5), (0.5))])

    def __getitem__(self, idx):
        x = self.dataset[idx]
        return self.transform(x)
        
    def __len__(self):
        return len(self.dataset)

DataLoader初始化:

train_set = MyDataset(data, image_size=image_size)
data_loader = DataLoader(train_set, batch_size=batch_size, shuffle=True)

但调用时发现,next(iter(data_loader))或者循环中的data,如果取data[0]得到的是单个样本,而非整批数据。移除[0]后代码能正常运行,但我担心遗漏了关键的数据形状处理步骤,怕后续出现错误。

问题原因与说明

这是因为官方的ImageFolder返回的是**(图像张量, 标签)**的元组,所以DataLoader会将每个batch打包成(批量图像张量, 批量标签张量)的结构,因此需要用data[0]提取图像部分。

而你的自定义MyDataset的__getitem__方法只返回了图像张量,没有返回标签,所以DataLoader直接将batch打包成批量图像张量,不需要再取[0]。这种情况完全正常,不会引发后续错误,只要确保:

  • 张量形状符合模型输入要求:你的单通道图像张量形状应为(batch_size, 1, image_size, image_size),而官方是3通道的(batch_size, 3, image_size, image_size),后续生成器和判别器的输入通道数需要对应修改为1。
  • 可以通过打印data.shape确认批量张量的结构是否正确。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 00:06:25