代码之家  ›  专栏  ›  技术社区  ›  davidwangv5

是torchvision.datasets.cifar。CIFAR10是否有列表?

  •  1
  • davidwangv5  · 技术社区  · 9 年前

    trainset = torchvision.datasets.CIFAR10(root='./data', train=True,
                                        download=True, transform=transform)
    print(trainset[1])
    print(trainset[:10])
    print(type(trainset))
    

    然而,我在尝试时遇到了一些错误

    print(trainset[:10])
    

    错误信息为

    TypeError: Cannot handle this data type
    

    我想知道为什么我可以使用 trainset[1] ,但不是 trainset[:10]

    2 回复  |  直到 5 年前
        1
  •  2
  •   Community Mohan Dere    6 年前

    CIFAR10不支持切片,这就是为什么会出现这种错误。如果你想要前10个,你必须这样做:

    print([trainset[i] for i in range(10)])
    

    更多信息

    可以索引CIFAR10类实例的主要原因是该类实现了 __getitem__() 作用

    trainset[i] 你实际上是在打电话 trainset.__getitem__(i)

    现在,在python3中,切片表达式也通过 其中切片表达式传递给 __getitem__() 作为切片对象。

    trainset[2:10] 相当于 trainset.__getitem__(slice(2, 10))

    由于两种不同类型的对象被传递到 __getitem__ 被期望做完全不同的事情,你必须明确地处理它们。

    不幸的是,正如你从 CIFAR10类的方法实现:

    def __getitem__(self, index):
        if self.train:
            img, target = self.train_data[index], self.train_labels[index]
        else:
            img, target = self.test_data[index], self.test_labels[index]
    
        # doing this so that it is consistent with all other datasets
        # to return a PIL Image
        img = Image.fromarray(img)
    
        if self.transform is not None:
            img = self.transform(img)
    
        if self.target_transform is not None:
            target = self.target_transform(target)
    
        return img, target
    
        2
  •  0
  •   tschomacker    5 年前

    https://stackoverflow.com/a/45226879/7924573 entrophys答案我建议使用 torch.utils.data.dataset.random\u split e、 g.这种方式:

    train_size = int(0.8*len(dataset))
    test_size = len(dataset) - train_size
    lengths = [train_size, test_size]
    train_dataset, valid_dataset = torch.utils.data.dataset.random_split(dataset, lengths)
    trainloader = DataLoader(train_data, 
      batch_size=args.train_batch,  
      shuffle=True, 
      num_workers=args.nThreads, 
      pin_memory=True)
    validloader = DataLoader(valid_data, 
      batch_size=args.train_batch,  
      shuffle=True, 
      num_workers=args.nThreads, 
      pin_memory=True)
    

    资料来源: https://yimjiyoung.github.io/2020/02/13/How-to-split-dataset-into-train-and-validation-set-in-pytorch/