读取TensorFlow TFRecord时解析sign字段失败,求助解决方案
解决TFRecord解析时的InvalidArgumentError问题
你的报错根源是TFRecord写入和读取时的特征类型不匹配:
- 写入TFRecord时,
sign字段使用的是tf.train.Feature(对应单值或一维列表的Int64List) - 读取时却错误使用了
tf.io.FixedLenSequenceFeature,这个API用于解析tf.train.FeatureList(多维度序列数据),两者无法兼容,导致解析失败。
根据你写入时的逻辑,提供两种针对性修改方案:
方案一:如果phrase的长度固定(比如固定为59)
修改decode_tfrec函数中sign的特征定义,改用FixedLenFeature并指定固定长度:
def decode_tfrec(record_bytes): features = tf.io.parse_single_example(record_bytes, { 'coordinates': tf.io.FixedLenFeature([], tf.string), # 替换FixedLenSequenceFeature为FixedLenFeature,指定固定长度shape 'sign': tf.io.FixedLenFeature([59], tf.int64), }) out = {} out['coordinates'] = tf.reshape(tf.io.decode_raw(features['coordinates'], tf.float32), (-1,ROWS_PER_FRAME,3)) out['sign'] = features['sign'] return out
方案二:如果phrase的长度不固定
改用VarLenFeature解析可变长度的一维整数列表,再转为密集张量:
def decode_tfrec(record_bytes): features = tf.io.parse_single_example(record_bytes, { 'coordinates': tf.io.FixedLenFeature([], tf.string), # 用VarLenFeature解析可变长度的int64列表 'sign': tf.io.VarLenFeature(tf.int64), }) out = {} out['coordinates'] = tf.reshape(tf.io.decode_raw(features['coordinates'], tf.float32), (-1,ROWS_PER_FRAME,3)) # 将稀疏张量转为密集张量 out['sign'] = tf.sparse.to_dense(features['sign']) return out
额外验证点
确保写入TFRecord时的phrase是一维整数列表,如果是二维数组,需要调整写入逻辑(改用FeatureList),但从你的代码看,phrase应为一维结构,上述方案即可解决问题。
内容的提问来源于stack exchange,提问作者Lệ Kiệt
相关产品推荐
相关产品推荐

