【发布时间】: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 永远不会终止。