【问题标题】:Batches of points with the same label on PytorchPytorch 上具有相同标签的点的批次
【发布时间】:2020-06-28 17:20:20
【问题描述】:

我想在每个包含 N 个训练点的批次上使用梯度下降来训练神经网络。我希望这些批次只包含具有相同标签的点,而不是从训练集中随机抽样。

例如,如果我正在使用 MNIST 进行训练,我希望得到如下所示的批次:

batch_1 = {0,0,0,0,0,0,0,0}

batch_2 = {3,3,3,3,3,3,3,3}

batch_3 = {7,7,7,7,7,7,7,7}

.....

等等。

我如何使用 pytorch 来做到这一点?

【问题讨论】:

  • 每个类的点数不同,不需要被batch_size整除。那么你将如何处理呢?是否应该有一些包含不同类别的批次(例如,在某些时候,您将剩下 3 个等级 0 的点)或者您是否想要删除不适合批次的点?
  • 训练点的数量可能也不能被batch_size整除,所以应该没问题吧?
  • 例如 0 类有 5923 个点,所以如果你把它们分成大小为 8 的批次,你将有 740 个这样的批次(740*8 = 5920),将有 3 个点还剩 0 个。你把它们放在哪里?
  • 当您有 50.000 个训练点和 128 个批量大小时会发生什么?它们也不是可分的,但这是一个非常常见的设置。为了回答您的问题,我可以在特定时期放弃一些分数。谢谢!
  • 我对 QMNIST 的看法有误,对此感到抱歉,删除了我的答案。

标签: python classification pytorch


【解决方案1】:

一种方法是为每个类创建子集和数据加载器,然后通过在每次迭代时在数据加载器之间随机切换来进行迭代:

import torch
from torch.utils.data import DataLoader, Subset
from torchvision.datasets import MNIST
from torchvision import transforms
import numpy as np

dataset = MNIST('path/to/mnist_root/', 
                transform=transforms.ToTensor(),
                download=True)

class_inds = [torch.where(dataset.targets == class_idx)[0]
              for class_idx in dataset.class_to_idx.values()]

dataloaders = [
    DataLoader(
        dataset=Subset(dataset, inds),
        batch_size=8,
        shuffle=True,
        drop_last=False)
    for inds in class_inds]

epochs = 1

for epoch in range(epochs):
    iterators = list(map(iter, dataloaders))   
    while iterators:         
        iterator = np.random.choice(iterators)
        try:
            images, labels = next(iterator)   
            print(labels)
            # do_more_stuff()

        except StopIteration:
            iterators.remove(iterator)

这适用于任何数据集(不仅仅是 MNIST)。 这是每次迭代打印标签的结果:

tensor([6, 6, 6, 6, 6, 6, 6, 6])
tensor([3, 3, 3, 3, 3, 3, 3, 3])
tensor([0, 0, 0, 0, 0, 0, 0, 0])
tensor([5, 5, 5, 5, 5, 5, 5, 5])
tensor([8, 8, 8, 8, 8, 8, 8, 8])
tensor([0, 0, 0, 0, 0, 0, 0, 0])
...
tensor([1, 1, 1, 1, 1, 1, 1, 1])
tensor([1, 1, 1, 1, 1, 1])

请注意,通过设置drop_last=False,会有批次,这里和那里,少于batch_size 元素。通过将其设置为 True,批次的大小将全部相同,但会丢弃一些数据点。

【讨论】:

    猜你喜欢
    • 2022-12-11
    • 1970-01-01
    • 2015-07-21
    • 1970-01-01
    • 2021-11-28
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多