【问题标题】:How can one clone a Pytoch dataset into another variable?如何将 Pytorch 数据集克隆到另一个变量中?
【发布时间】:2020-02-10 20:23:27
【问题描述】:

我想创建 Pytorch 中提供的 MNIST 数据集的几个子集。每个子集应该有不同的类。我尝试的是以下内容:

def split_MNIST(mnist_set, digits):
    dset = mnist_set
    classes = []
    indices = dset.targets == digits[0]
    classes.append(dset.classes[digits[0]])
    if len(digits) > 1:
        for digit in digits[1:]:
            idx = dset.targets == digit
            indices = indices + idx
            classes.append(dset.classes[digit])
    dset.targets = dset.targets[indices]
    dset.data = dset.data[indices]
    dset.classes = classes
    return dset


train = datasets.MNIST("../data", train=True, download=True,
                        transform=transforms.Compose([transforms.ToTensor()]))

test =datasets.MNIST("../data", train=False, download=True,
                      transform=transforms.Compose([transforms.ToTensor()]))

tr = split_MNIST(train, [1,2,3])

trainset = torch.utils.data.DataLoader(tr, batch_size=16, shuffle=True)

这行得通,但它实际上改变了原始训练变量,而不是创建新数据集。有没有办法创建数据集的克隆而不是保留原始数据集?

【问题讨论】:

  • 可能最直接的方法是使用 torch.utils.data.Subset 并提供所需样本的索引。每个子集都保留对原始数据集的引用,并对其索引列表中提供的相应元素进行采样。
  • 你可以取对象ex的copy.deepcopy。在你的情况下dset = copy.deepcopy(mnist_set) 可以正常工作

标签: python-3.x pytorch sampling


【解决方案1】:

只需将数据集实例化放在split_MNIST func 中即可。

def split_MNIST(path2data, train, download, transform, digits):
    dset = datasets.MNIST(path2data, train=train, download=download, transform=transform)
    classes = []
    indices = dset.targets == digits[0]
    classes.append(dset.classes[digits[0]])
    if len(digits) > 1:
        for digit in digits[1:]:
            idx = dset.targets == digit
            indices = indices + idx
            classes.append(dset.classes[digit])
    dset.targets = dset.targets[indices]
    dset.data = dset.data[indices]
    dset.classes = classes
    return dset


transforms = transforms.Compose([transforms.ToTensor()])
tr = split_MNIST('../data', train=True, download=True, transform=transforms, digits=[1,2,3])

trainset = torch.utils.data.DataLoader(tr, batch_size=16, shuffle=True)

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2018-12-13
    • 1970-01-01
    • 1970-01-01
    • 2016-09-08
    • 1970-01-01
    • 2014-04-29
    • 1970-01-01
    • 2011-04-29
    相关资源
    最近更新 更多