如何仅保留TFRecordDataset首条记录并删除其余所有数据
问题说明
- 基于
test_filenames创建的TFRecordDataset共包含10000条记录,初始化代码如下:
test_dataset = tf.data.TFRecordDataset([test_filenames])
- 需求为仅保留数据集第一条记录、删除其余记录,预期功能伪代码:
test_dataset = test_dataset.removeAllExceptFirst() ... first_record = test_dataset.getItem(0) test_dataset = test_dataset.removeAll() test_dataset = test_dataset.add(first_record)
- 此前尝试使用
test_dataset.batch(1).take(1)方案实现,运行触发报错,测试代码:
def test_function(record): keys_to_features = { "test1": tf.io.FixedLenFeature((), tf.string, default_value=""), 'test2': tf.io.FixedLenFeature([], tf.string), "test3": tf.io.FixedLenFeature((), tf.string) } features = tf.io.parse_single_example(record, keys_to_features) print("features: {}".format(features)) return None, None test_dataset = tf.data.TFRecordDataset([test_filenames]) test_dataset = test_dataset.batch(1).take(1) test_dataset = test_dataset.map(test_function)
- 核心报错信息:
ValueError: Input serialized must be a scalar
解决方法
报错原因
调用batch(1)后,数据集输出的每个元素不再是单条序列化Example标量,而是形状为(1,)的批量张量,tf.io.parse_single_example要求输入必须是标量,因此触发类型错误。
正确实现方式
TensorFlow 内置take(n)方法可以直接实现「仅保留前n条记录、丢弃其余记录」的需求,要取第一条直接调用take(1)即可,不需要额外套batch(1)。
注意tf.data.Dataset是惰性迭代的数据集对象,不支持Python列表式的随机索引访问、原地增删操作,伪代码里的removeAllExceptFirst、getItem、removeAll、add都不是Dataset的内置API,所有数据变换都要通过Dataset提供的转换算子生成新数据集实现。
两种可运行的正确代码写法:
- 常规流水线写法(推荐):先做单样本解析,再取第一条
# 初始化数据集 test_dataset = tf.data.TFRecordDataset([test_filenames]) # 定义单样本解析函数 def parse_function(record): keys_to_features = { "test1": tf.io.FixedLenFeature((), tf.string, default_value=""), 'test2': tf.io.FixedLenFeature([], tf.string), "test3": tf.io.FixedLenFeature((), tf.string) } features = tf.io.parse_single_example(record, keys_to_features) return features # 先解析单条数据,再取第一条 test_dataset = test_dataset.map(parse_function).take(1)
- 批量处理写法:先取第一条再组batch,用批量解析接口适配
test_dataset = tf.data.TFRecordDataset([test_filenames]) # 先取第一条再组batch,避免处理多余数据 test_dataset = test_dataset.take(1).batch(1) def parse_batch_function(records): keys_to_features = { "test1": tf.io.FixedLenFeature((), tf.string, default_value=""), 'test2': tf.io.FixedLenFeature([], tf.string), "test3": tf.io.FixedLenFeature((), tf.string) } # 批量输入对应使用parse_example接口 features = tf.io.parse_example(records, keys_to_features) return features test_dataset = test_dataset.map(parse_batch_function)
如果需要提取第一条记录的具体数值,不需要实现getItem(0)方法,直接通过迭代器取首个元素即可:
# 简洁写法 first_record = next(iter(test_dataset)) # 等价遍历写法 for record in test_dataset: first_record = record break
内容的提问来源于stack exchange,提问作者stackbiz
相关产品推荐
相关产品推荐

