【问题标题】:Custom data loader is returning list in pytorch自定义数据加载器在 pytorch 中返回列表
【发布时间】:2020-06-17 06:15:56
【问题描述】:

我想从 3 个不同的文件夹中获取 3 批图像。我在 pytorch 中编写了自定义数据加载器。但它返回的列表一次包含所有批次而不是单个批次。(在 google colab 中运行)

#custom data loader
class set(Dataset):
    def __init__(self, dataset_input, dataset_expertA, dataset_expertB):
        self.dataset1 = dataset_input
        self.dataset2 = dataset_expertA
        self.dataset3 = dataset_expertB

    def __getitem__(self, index):
        x1 = self.dataset1[index]
        x2 = self.dataset2[index]
        x3 = self.dataset3[index]

        return x1, x2, x3

    def __len__(self):
        return len(self.dataset1)

input_path = "/content/gdrive/My Drive/project/input/"

dataset = datasets.ImageFolder(root= input_path, transform=transforms.Compose([
                               transforms.Resize([64,64]),
                               transforms.ToTensor(),
                               transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
                               ]))

expertA_path = "/content/gdrive/My Drive/project/expertA/"

datasetA = datasets.ImageFolder(root= expertA_path, transform=transforms.Compose([
                               transforms.Resize([64,64]),
                               transforms.ToTensor(),
                               transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
                               ]))


expertB_path = "/content/gdrive/My Drive/project/expertB/"

datasetB = datasets.ImageFolder(root= expertB_path, transform=transforms.Compose([
                               transforms.Resize([64,64]),
                               transforms.ToTensor(),
                               transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
                               ]))


data = set(dataset, datasetA, datasetB)
dataloader = torch.utils.data.DataLoader(data, batch_size=64,
                                         shuffle=True, num_workers=2)


for i, (inp, expA, expB) in enumerate(dataloader):

  print(inp.shape)
  break

这会打印出 inp 是列表的错误,当我 print(inp[0].shape) 我得到正确的形状时,我认为 inp 包含所有批次,即 inp[0]、inp[1]...

我在数据加载器代码中犯了什么错误?

【问题讨论】:

    标签: python pytorch


    【解决方案1】:

    datasets.ImageFolder 返回一个 (image, label) 的元组,因此inp 也是一个元组,其中inp[0] 是图像,inp[1] 是它们对应的标签。这同样适用于expAexpB

    如果您只想要没有标签的图像,您可以忽略标签并在访问自定义数据集中的数据时只返回图像:

    def __getitem__(self, index):
        image1, label1 = self.dataset1[index]
        image2, label2 = self.dataset2[index]
        image3, label3 = self.dataset3[index]
    
        return image1, image2, image3
    

    【讨论】:

      猜你喜欢
      • 2021-11-09
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2018-11-25
      • 2021-12-29
      • 2017-09-12
      • 1970-01-01
      • 1970-01-01
      相关资源
      最近更新 更多