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

如何使用tf.train.Feature与tf.train.Example存储Ragged Tensor到TFRecords

解决方案核心思路

tf.train.Example 仅支持存储定长一维列表,无法直接写入不规则结构的Ragged Tensor,需要将其拆分为两个独立字段存储,完全兼容常规张量的写入逻辑:

  • 所有元素展开后的扁平化一维数组
  • 对应维度的行分割偏移量数组,用来还原Ragged结构
写入TFRecords实现代码
import tensorflow as tf

# 测试用Ragged Tensor
d = [[], [1, 3], [5]]
d_tf = tf.ragged.constant(d)

# 提取Ragged Tensor的核心存储信息
# 1. 扁平化后的所有元素值
flat_vals = d_tf.flat_values.numpy()
# 2. 行分割偏移量(一维Ragged仅需要一层row_splits,多层嵌套按需存多层即可)
row_splits = d_tf.row_splits.numpy()

# 序列化逻辑,支持同时写入常规张量
def serialize_example(ragged_flat, ragged_splits, regular_feat):
    feature = {
        # Ragged Tensor对应的两个存储字段
        "ragged_flat": tf.train.Feature(int64_list=tf.train.Int64List(value=ragged_flat)),
        "ragged_row_splits": tf.train.Feature(int64_list=tf.train.Int64List(value=ragged_splits)),
        # 常规张量正常写入即可
        "regular_feat": tf.train.Feature(int64_list=tf.train.Int64List(value=[regular_feat]))
    }
    example = tf.train.Example(features=tf.train.Features(feature=feature))
    return example.SerializeToString()

# 测试写入
with tf.io.TFRecordWriter("ragged_test.tfrecord") as writer:
    # 示例常规特征取值为200
    serialized_str = serialize_example(flat_vals, row_splits, 200)
    writer.write(serialized_str)
读取并还原Ragged Tensor代码
# 定义样本解析函数
def parse_example(example_proto):
    feature_desc = {
        "ragged_flat": tf.io.VarLenFeature(tf.int64),
        "ragged_row_splits": tf.io.VarLenFeature(tf.int64),
        "regular_feat": tf.io.FixedLenFeature([], tf.int64)
    }
    parsed = tf.io.parse_single_example(example_proto, feature_desc)
    # 转换稀疏特征为稠密Tensor
    ragged_flat = tf.sparse.to_dense(parsed["ragged_flat"])
    ragged_splits = tf.sparse.to_dense(parsed["ragged_row_splits"])
    # 还原Ragged Tensor结构
    restored_ragged = tf.RaggedTensor.from_row_splits(ragged_flat, ragged_splits)
    return restored_ragged, parsed["regular_feat"]

# 测试读取
dataset = tf.data.TFRecordDataset("ragged_test.tfrecord")
dataset = dataset.map(parse_example)
for ragged_feat, regular_val in dataset:
    print("还原后Ragged Tensor:", ragged_feat)
    print("常规特征取值:", regular_val.numpy())
注意事项
  • 多层嵌套的高维Ragged Tensor,只需要将每一层的row_splits单独存为一个特征字段,还原时按层级拼接即可
  • 存储浮点型Ragged Tensor时,将代码中的tf.train.Int64List替换为tf.train.FloatList即可,核心逻辑完全一致

内容的提问来源于stack exchange,提问作者eng2019

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 20:36:05