【发布时间】:2020-12-16 15:31:41
【问题描述】:
我有一个数据集,当我使用ds = ds.map(process_path, num_parallel_calls=AUTOTUNE).prefetch(AUTOTUNE) 对其进行预处理时,该行的执行速度非常快。然后,当我尝试使用以下方法访问数据集的元素之一时:
for image, label in ds.take(1):
print(image.shape)
image = tf.squeeze(image)
plt.imshow(image, cmap='gray')
加载需要一两秒钟;这是我的第一个问题:
预处理是否仅在访问数据集中的元素时在数据集上运行,而不是在我调用 ds.map(process_path,...) 时立即运行?
但是我的主要问题是,当我将数据集 ds 分成两部分(训练和测试)并尝试再次访问其中一个元素时,速度相当慢......就像慢了 20 倍。我把它分成两部分:
test_ds_size = int(image_count * 0.2)
train_ds = ds.skip(test_ds_size)
test_ds = ds.take(test_ds_size)
然后我尝试以与上述相同的方式访问它,但将 ds 替换为 train_ds;我的第二个问题是:
为什么这会慢得多,只是将它分成两部分?
还是我做错了什么……
非常感谢任何帮助。
【问题讨论】:
-
如果没有其他数据集管道,就很难诊断。不过,您的第一个直觉是正确的:
dataset.map不会加载任何数据。
标签: python tensorflow machine-learning tensorflow2.0