如何从TensorFlow的tf.data.TFRecordDataset取指定索引区间元素
TensorFlow Dataset 按索引范围取元素实现
问题场景
按如下方式定义test_dataset:
test_dataset = tf.data.TFRecordDataset([test_tfrecords]) test_dataset = test_dataset.map(map_f) test_dataset = test_dataset.repeat(1) test_dataset = test_dataset.batch(1)
常规获取前100个元素的写法是用take()方法:
for test in test_dataset.take(100): pass
如果要获取index range(索引范围)为50到150之间的元素,直接给take()传入区间列表的写法是无效的,运行达不到预期效果:
# 错误写法,不支持传入区间列表参数 for test in test_dataset.take([50-150]): pass
正确实现方案
tf.data.Dataset 不支持直接给take()传索引区间切片,搭配skip()方法就能实现范围取数:skip(n)的作用是跳过数据集开头的n个元素,跳过之后的第一个元素对应的原始索引就是n,后续再接take()取指定数量的元素,就能精准拿到目标区间的内容。
- 如果你需要取索引为50到149的元素(跳过前50个索引为0-49的元素,取后续连续100个),写法如下:
for test in test_dataset.skip(50).take(100): # 写入你的元素处理逻辑 pass
- 如果你需要把索引为150的元素也包含进来(完整覆盖50到150的索引范围,共101个元素),只需要把
take()的参数调整为101即可:
for test in test_dataset.skip(50).take(101): pass
注意:定义数据集时调用了
repeat(1),意味着数据集只会完整迭代一轮,只要skip和take的参数之和不超过数据集本身的总样本量就能正常取数;如果参数总和超出数据集实际大小,迭代到数据集末尾时会自动停止,不会抛出报错。
内容的提问来源于stack exchange,提问作者stackbiz
相关产品推荐
相关产品推荐

