【问题标题】:Preventing PyTorch Dataset iteration from exceeding length of dataset防止 PyTorch 数据集迭代超过数据集长度
【发布时间】:2019-09-02 08:54:50
【问题描述】:

我正在使用具有以下内容的自定义 PyTorch 数据集:

class ImageDataset(Dataset):
    def __init__(self, input_dir, input_num, input_format, transform=None):
        self.input_num = input_num
        # etc
    def __len__ (self):
        return self.input_num
    def __getitem__(self,idx):
        targetnum = idx % self.input_num
        # etc

但是,当我迭代此数据集时,迭代会循环回到数据集的开头,而不是在数据集的末尾终止。这实际上变成了迭代器中的无限循环,epoch print 语句永远不会出现在后续的 epoch 中。

train_dataset=ImageDataset(input_dir = 'path/to/directory', 
                           input_num = 300, input_format = "mask") # Size 300
num_epochs = 10
for epoch in range(num_epochs):
    print("EPOCH " + str(epoch+1) + "\n")
    num = 0
    for data in train_dataset:
        print(num, end=" ")
        num += 1
        # etc

打印输出(...对于介于两者之间的值):

EPOCH 1
0 1 2 3 4 5 6 7 ... 298 299 300 301 302 303 304 305 ... 597 598 599 600 601 602 603 604 ...

为什么对 Dataset 的基本迭代继续超过 DataSet 的定义 __len__,以及如何确保在使用此方法时达到数据集长度后对数据集的迭代终止(或手动迭代数据集长度的范围是唯一的解决方案)?

谢谢。

【问题讨论】:

  • 为什么不使用 DataLoader? train_loader = torch.utils.data.DataLoader(train_dataset)?它就是为此目的而创建的。
  • 使用 DataLoader 可能是最好的,但我仍然不明白为什么迭代 DataSet 永远不会终止。

标签: python pytorch


【解决方案1】:

Dataset 类没有实现StopIteration 信号。

for 循环监听 StopIteration。 for 语句的目的是循环遍历迭代器提供的序列,异常用于表示迭代器现在已经完成...

更多:Why does next raise a 'StopIteration', but 'for' do a normal return? | The Iterator Protocol

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2020-06-21
    • 2016-10-27
    • 2020-05-23
    • 2012-04-23
    • 2019-07-22
    • 2022-07-18
    • 2021-12-15
    • 2021-12-28
    相关资源
    最近更新 更多