如何跳过开头元素提取tf.data.Dataset中的指定数据项
访问tf.data.Dataset非起始位置内容的核心方法是链式组合skip()和take()方法:skip(n)会跳过数据集前n个计数单位(单位根据数据集预处理状态决定:未调用batch()时是单个样本,调用batch()分批后是单个批次),返回剩余内容的数据集视图,后续衔接take(m)即可从剩余内容的开头提取m个单位的内容。两个方法均为惰性执行,不会额外加载不需要的数据,性能和直接提取数据集开头元素无明显差异。
1. 指定位置单个元素提取
仅修改take()参数为2仍只能拿到第一个元素,是因为原有循环在第一次迭代拿到第一个元素后就触发break终止,不会读取第二个元素。提取指定位置单个元素时,先跳过目标位置前的所有元素,再取1个元素即可(注意数据集索引从0开始计数,第k个元素对应需要跳过k-1个前置元素)。
示例代码(提取测试集第二个元素):
# 目标为第2个元素,索引为1,因此跳过前1个元素,提取1个元素 target_index = 1 for image, label in test_dataset.skip(target_index).take(1): break image = image.numpy().reshape((28,28))
如果需要提取第5个元素,仅需将target_index修改为4即可。
2. 指定起始位置批量图像可视化
需求为从第100个元素开始读取25张图像,第100个元素对应的索引为99,因此跳过前99个元素后,连续提取25个元素即可实现需求。
修改后的完整代码:
plt.figure(figsize=(10,10)) # 跳过前99个元素(对应从第100个元素开始),连续提取25个元素 for i, (image, label) in enumerate(train_dataset.skip(99).take(25)): image = image.numpy().reshape((28,28)) plt.subplot(5,5,i+1) plt.xticks([]) plt.yticks([]) plt.grid(False) plt.imshow(image, cmap=plt.cm.binary) plt.xlabel(class_names[label]) plt.show()
如果需要调整起始位置,比如从第200个元素开始读取,仅需将skip()的参数修改为199即可,take()的参数可根据需要提取的图像数量灵活调整。
3. 指定批次数据读取预测
如果数据集已经调用batch()方法完成分批,skip()和take()的计数单位会自动变为批次而非单个样本:要读取第n个批次的数据,仅需跳过前n-1个批次,再提取1个批次即可。
示例代码(读取第二个批次执行预测):
# 目标为第2个批次,索引为1,因此跳过前1个批次,提取1个批次 target_batch_index = 1 for test_images, test_labels in test_dataset.skip(target_batch_index).take(1): test_images = test_images.numpy() test_labels = test_labels.numpy() predictions = model.predict(test_images)
如果需要读取第三个批次,仅需将target_batch_index修改为2即可。
注意:如果数据集在分批前调用过
shuffle()做随机打乱,需要提前固定随机种子,否则每次运行时数据集顺序会随机变化,无法稳定获取固定位置的内容。
内容的提问来源于stack exchange,提问作者Giorgio Tempesta

