如何将TFRecord字节串解析为张量字典?多任务Transformer适配
解决TFRecords可变长度张量解析问题
核心思路
针对你样本内张量长度一致、样本间长度可变的结构,且写入时已将序列张量转为字节串存储的情况,解析关键是用tf.io.parse_tensor将字节串还原为原类型张量,再重组为目标字典结构。
完整可复现代码
1. 生成模拟数据并写入TFRecords
import numpy as np import tensorflow as tf # 生成模拟样本:样本间序列长度随机可变 def generate_sample(seq_len): return { 'continuous_input': tf.random.normal((seq_len,), dtype=tf.float32), 'categorical_input': tf.random.uniform((seq_len,), minval=0, maxval=10, dtype=tf.int32), 'continuous_output': tf.random.normal((seq_len,), dtype=tf.float32), 'categorical_output': tf.random.uniform((seq_len,), minval=0, maxval=5, dtype=tf.int32) } # 写入TFRecords def write_tfrecords(file_path, num_samples): with tf.io.TFRecordWriter(file_path) as writer: for _ in range(num_samples): seq_len = np.random.randint(5, 20) # 模拟样本间长度差异 sample = generate_sample(seq_len) # 将每个张量序列转为字节串存储 feature = { 'continuous_input': tf.train.Feature(bytes_list=tf.train.BytesList(value=[tf.io.serialize_tensor(sample['continuous_input']).numpy()])), 'categorical_input': tf.train.Feature(bytes_list=tf.train.BytesList(value=[tf.io.serialize_tensor(sample['categorical_input']).numpy()])), 'continuous_output': tf.train.Feature(bytes_list=tf.train.BytesList(value=[tf.io.serialize_tensor(sample['continuous_output']).numpy()])), 'categorical_output': tf.train.Feature(bytes_list=tf.train.BytesList(value=[tf.io.serialize_tensor(sample['categorical_output']).numpy()])) } example = tf.train.Example(features=tf.train.Features(feature=feature)) writer.write(example.SerializeToString()) # 生成示例数据文件 write_tfrecords('multi_task_data.tfrecord', 10)
2. 修正后的解析函数与读取流程
# 解析TFRecords核心函数 def parse_tfrecord_fn(example_proto): # 定义特征描述:所有特征均为单个字节串 feature_description = { 'continuous_input': tf.io.FixedLenFeature([], tf.string), 'categorical_input': tf.io.FixedLenFeature([], tf.string), 'continuous_output': tf.io.FixedLenFeature([], tf.string), 'categorical_output': tf.io.FixedLenFeature([], tf.string) } # 解析单个example的字节串特征 parsed_features = tf.io.parse_single_example(example_proto, feature_description) # 将字节串还原为对应类型的张量 parsed_sample = { 'continuous_input': tf.io.parse_tensor(parsed_features['continuous_input'], out_type=tf.float32), 'categorical_input': tf.io.parse_tensor(parsed_features['categorical_input'], out_type=tf.int32), 'continuous_output': tf.io.parse_tensor(parsed_features['continuous_output'], out_type=tf.float32), 'categorical_output': tf.io.parse_tensor(parsed_features['categorical_output'], out_type=tf.int32) } return parsed_sample # 构建数据集并验证解析结果 dataset = tf.data.TFRecordDataset('multi_task_data.tfrecord') dataset = dataset.map(parse_tfrecord_fn) # 打印第一个样本验证解析效果 for sample in dataset.take(1): print("解析后的样本结构:") for key, tensor in sample.items(): print(f"{key}: 数据类型={tensor.dtype}, 序列长度={tensor.shape[0]}")
关键细节说明
- 写入阶段:用
tf.io.serialize_tensor将可变长度张量转为字节串,确保序列长度信息被完整保留。 - 解析阶段:
- 先用
tf.io.parse_single_example读取字节串特征,特征描述使用tf.io.FixedLenFeature([], tf.string),因为每个特征对应单个字节串。 - 再通过
tf.io.parse_tensor将字节串还原为原张量,指定out_type匹配原数据类型(float32/int32),自动恢复序列长度。
- 先用
- 该方案完美适配样本间长度可变的场景,完全还原原始张量的形状与类型。
内容的提问来源于stack exchange,提问作者David Bellamy
相关产品推荐
相关产品推荐

