如何设置TensorFlow中ParallelMapDataset类型数据集的图片数量
错误原因
tf.data.Dataset及其子类(包括你遇到的ParallelMapDataset)不支持Python原生的下标索引语法,因此使用test_images[100]会触发不可下标访问的报错。
正确实现方案
如果要截取测试集的前100张图片,调用数据集内置的take()方法即可:
test_images = dataset['test'].take(100)
后续构建test_batches的代码不需要做任何修改,take()返回的仍然是符合tf.data规范的数据集对象,和batch、prefetch等操作完全兼容。
如果需要随机抽取100张而非按顺序取前100张,可以先对测试集做打散再截取:
# 固定seed保证抽取结果可复现 test_images = dataset['test'].shuffle( buffer_size=info.splits['test'].num_examples, seed=42 ).take(100)
补充常用操作
- 如需跳过指定数量的样本再截取,可搭配
skip()方法使用,比如dataset['test'].skip(200).take(100)会跳过前200张,取第201到300张样本。 - 所有操作返回的都是
tf.data.Dataset类型,无需额外类型转换即可接入后续的数据处理流程。
内容的提问来源于stack exchange,提问作者Vikram
相关产品推荐
相关产品推荐

