You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.01 13:09:04