【问题标题】:How is data augmentation done in each epoch?每个时期如何进行数据扩充?
【发布时间】:2020-10-10 01:53:18
【问题描述】:

我是 PyTorch 的新手,想对每个 epoch 的数据集应用数据扩充。我

train_transform = Compose([
    transforms.RandomHorizontalFlip(p=0.5),
    transforms.RandomCrop(32, padding=4),
    transforms.ToTensor(),
    transforms.Normalize([0, 0, 0], [1, 1, 1])
])

test_transform = Compose([
    transforms.ToTensor(),
    transforms.Normalize([0, 0, 0], [1, 1, 1])
])

cifar10_train = CIFAR10(root = "/data", train=True, download = True, transform=train_transform)
train_loader = torch.utils.data.DataLoader(cifar10_train, batch_size=128, shuffle=True)

cifar10_test = CIFAR10(root = "/data", train=False, download = True, transform=test_transform)
test_loader = torch.utils.data.DataLoader(cifar10_test, batch_size=128, shuffle=True)

我从在线教程中获得了代码。因此,据我了解 train_transform 和 test_transform 是增强代码,而 cifar10_train 和 cifar10_test 是加载数据并同时完成增强的地方。这是否意味着数据增强只在训练前进行一次?如果我想为每个 epoch 做数据增强怎么办。

【问题讨论】:

    标签: python pytorch data-augmentation


    【解决方案1】:

    我认为您的代码存在一些误解。 cifar10_traincifar10_test 实际上将数据集加载到 python 中(此数据未增强,是原始数据),然后数据经过转换。在大多数情况下,训练集是完成数据扩充的地方,而测试集没有扩充,因为它应该复制真实世界的数据。转换(train_transformtest_transforms)决定了如何对数据进行扩充、规范化和转换为 PyTorch 张量,您可以将其视为数据集要遵循的一组准则/规则。如前所述,训练集只会被增强,这就是为什么 train_transform 有 RandomHorizontalFlipRandomCrop(它进行增强),以及为什么 test_transforms 没有 RandomHorizontalFlipRandomCrop。加载器(train_loadertest_loader)将数据拆分为批次(数据组),并将转换应用于 cifar10 数据集。

    【讨论】:

    • 谢谢你的帮助,它帮助我理解了一些。但我仍然不明白的是。在每个 epoch 的训练过程中,训练数据总是会因为变换而具有不同的 shapeof 图像,对吧?因此,根据我在以下训练代码中看到的内容,对于 i in range(EPOCHS):start_time = time.time() ep = 0 model.train() for X_b, y_b in train_loader: optim.zero_grad() X_b = X_b.to(device) y_b = y_b.to(device) ```我没有看到他们在每个时期都应用转换。
    • 训练数据将始终具有相同形状的图像,神经网络只能接收恒定数量的数据,不能更小或更大。转换首先执行 RandomHorizo​​ntalFlip,然后将其调整为图像的设置高度和宽度。我认为您的另一个问题已经回答了,但再次回答,测试数据只转换一次,训练数据只转换一次。这通过设置 transfomrs=transforms 在CIFAR10 中发生。
    【解决方案2】:

    以下是粗略的操作流程:

    1. 从文件系统中读取 128 个示例。
    2. 批处理它们(即制作一批示例)并将转换应用于批处理。
    3. 将批次传递到网络。

    因此,在将其馈送到模型之前,会将转换应用于每个批次,而与时期无关。

    【讨论】:

    • 等等,数据集只会变换一次?我想如果我运行 10 个 epoch,对于每个 epoch,我都会得到不同版本的数据集。
    • 是的,这是正确的。这些步骤将在每个 epoch 中为每个批次执行 :)
    猜你喜欢
    • 2020-06-18
    • 2020-06-05
    • 2016-02-13
    • 2018-04-01
    • 2017-08-19
    • 1970-01-01
    • 2023-01-26
    • 1970-01-01
    • 2020-05-15
    相关资源
    最近更新 更多