如何使用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
相关产品推荐
相关产品推荐

