TensorFlow 2下如何将可变长度tuple/list/1-D array特征存入TFRecord
TFRecord存储可变长度序列特征的实现方法(TensorFlow 2)
核心原理
tf.train.Feature内置的Int64List/FloatList/BytesList本身就支持传入任意长度的1维可迭代对象(list、tuple、1D numpy数组都可以),不需要提前做填充,序列长度可以随单个样本动态变化,完全适配你提到的节点前驱/后继ID存储场景。
写入实现(以存储节点前驱/后继ID为例)
ID为整数类型,直接用Int64List存储即可:
import tensorflow as tf import numpy as np def serialize_example(node_feat, pre_ids, next_ids): feature = { # 原有固定长度浮点型节点特征 "node_feat": tf.train.Feature(float_list=tf.train.FloatList(value=node_feat)), # 可变长度前驱ID列表,pre_ids长度可以是0到任意数值,支持list/tuple/1D numpy数组 "pre_ids": tf.train.Feature(int64_list=tf.train.Int64List(value=pre_ids)), # 可变长度后继ID列表 "next_ids": tf.train.Feature(int64_list=tf.train.Int64List(value=next_ids)), } example_proto = tf.train.Example(features=tf.train.Features(feature=feature)) return example_proto.SerializeToString() # 不同长度序列的测试样例 # 样本1:2个前驱,3个后继 serialize_example([0.1, 0.2, 0.3], [101, 102], [201, 202, 203]) # 样本2:无可用前驱,1个后继 serialize_example([0.4, 0.5, 0.6], [], [204]) # 样本3:numpy数组格式的序列 serialize_example([0.7, 0.8, 0.9], np.array([103, 104, 105], dtype=np.int64), np.array([], dtype=np.int64))
注意:存储numpy数组格式的序列时,需要保证数据类型和对应List要求匹配:Int64List要求整数为int64类型,FloatList要求浮点数为float32/float64类型,类型不匹配时可以先调用
.astype(np.int64)做转换。
读取实现
读取时可变长度特征需要用tf.io.VarLenFeature声明类型,读出来默认是SparseTensor,可以根据需要转成DenseTensor或者保留稀疏格式直接参与运算:
# 特征格式声明 feature_description = { "node_feat": tf.io.FixedLenFeature([3], tf.float32), # 此处根据你的实际节点特征维度调整 "pre_ids": tf.io.VarLenFeature(tf.int64), # 可变长度特征声明 "next_ids": tf.io.VarLenFeature(tf.int64), } def parse_example(example_proto): return tf.io.parse_single_example(example_proto, feature_description) # 完整读取示例 def load_tfrecord(file_path): dataset = tf.data.TFRecordDataset(file_path) dataset = dataset.map(parse_example) for item in dataset: # 稀疏张量转稠密张量,不需要转的话可以省略这两步 item["pre_ids"] = tf.sparse.to_dense(item["pre_ids"]) item["next_ids"] = tf.sparse.to_dense(item["next_ids"]) yield item
其他类型可变长度序列适配
- 浮点型可变长度序列:写入用
FloatList,读取声明tf.io.VarLenFeature(tf.float32) - 字符串类型可变长度序列:写入用
BytesList(value=[s.encode('utf-8') for s in str_list]),读取声明tf.io.VarLenFeature(tf.string)
内容的提问来源于stack exchange,提问作者eng2019
相关产品推荐
相关产品推荐

