TensorFlow从TFRecord读取时序数据的Dataset结构疑问
首先,咱们先解决你遇到的ParseSingleSequenceExample报错问题,这个错误的核心原因是**FixedLenSequenceFeature的参数使用不正确**,再加上代码里的小语法错误。之后我会详细解释Dataset的结构,帮你理清怎么处理变长张量。
一、修复读取代码的错误
先看你的decode函数,这里有两个关键问题:
FixedLenSequenceFeature的第一个参数是序列中单个元素的形状,而不是整个序列的形状。你的feature_data是一维的浮点序列,每个元素就是单个float,所以shape应该是(),而不是(None,);同理labels是一维的整数序列,单个元素是int64,shape也是()。labels的定义里括号配对错误,你写成了((None,), tf.int64),多了一层括号。
修正后的decode函数应该是这样的:
def decode(serialized_proto): # 根据你实际写入的context补全特征,没有的话传空字典即可 context_features = {} sequence_features = { "feature_data": tf.FixedLenSequenceFeature((), tf.float32, allow_missing=True), "labels": tf.FixedLenSequenceFeature((), tf.int64, allow_missing=True) } context, sequence = tf.parse_single_sequence_example( serialized_proto, context_features=context_features, sequence_features=sequence_features ) return context, sequence
这里加上allow_missing=True是为了兼容可能的空序列(如果你的数据里有这种情况的话),如果每个SequenceExample都确保有这两个序列,可以去掉这个参数。
二、理解Dataset的结构与返回数据
当你用tf.data.TFRecordDataset读取TFRecord文件,再用map(decode)处理后,Dataset中的每个元素就是decode函数返回的元组(context, sequence):
context是一个字典,key是你定义的context特征名,value是对应形状的张量(比如如果context里有一个标量特征,那就是shape=[]的张量)。sequence也是一个字典,key是序列特征名(feature_data和labels),value是变长张量:比如feature_data的shape是[None],代表它是一维的变长序列,长度由每个SequenceExample中的实际数据决定;labels同理是[None]的int64张量。
遍历Dataset的方式
你之前用make_one_shot_iterator().get_next()是旧版本的写法,在TensorFlow 2.x中更推荐用迭代器直接遍历:
import tensorflow as tf dataset = tf.data.TFRecordDataset("data/tf_record.tfrecords") dataset = dataset.map(decode) # 遍历Dataset中的每个元素 for context, sequence in dataset: # 查看context内容(如果有定义的话) if context: print("Context features:", context) # 查看序列特征的形状和样本数据 print("feature_data shape:", sequence["feature_data"].shape) print("labels shape:", sequence["labels"].shape) # 取出具体数值(按需使用) print("feature_data sample:", sequence["feature_data"][:5].numpy()) print("labels sample:", sequence["labels"][:5].numpy())
如果要对接Estimator,需要把Dataset返回格式转换成Estimator要求的(features, labels)结构,同时注意变长序列的batch处理:
def input_fn(): dataset = tf.data.TFRecordDataset("data/tf_record.tfrecords") def decode_and_convert(serialized_proto): context, sequence = decode(serialized_proto) # 把sequence中的feature_data作为特征,labels作为标签 features = {"feature_data": sequence["feature_data"]} # 若需要包含context特征,可合并到features字典中 labels = sequence["labels"] return features, labels dataset = dataset.map(decode_and_convert) # 变长序列需用padded_batch做批量处理,自动填充到batch内最大长度 dataset = dataset.padded_batch( batch_size=32, padded_shapes={ "feature_data": [None], # None表示该维度长度可变 }, padding_values={ "feature_data": 0.0, # 指定填充的数值 } ) return dataset
三、补充写入代码的注意点
你的写入代码里np.random.random_integers已被numpy弃用,建议换成np.random.randint(0, 11, size=(1084,))(注意randint是左闭右开,写11才能取到0-10的整数)。另外,要记得序列化并写入后关闭writer:
writer.write(example.SerializeToString()) writer.close()
这样修改后,就能正常读取数据,也能清晰理解Dataset中每个元素的结构了。
内容的提问来源于stack exchange,提问作者e2718281

