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

如何将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将可变长度张量转为字节串,确保序列长度信息被完整保留。
  • 解析阶段:
    1. 先用tf.io.parse_single_example读取字节串特征,特征描述使用tf.io.FixedLenFeature([], tf.string),因为每个特征对应单个字节串。
    2. 再通过tf.io.parse_tensor将字节串还原为原张量,指定out_type匹配原数据类型(float32/int32),自动恢复序列长度。
  • 该方案完美适配样本间长度可变的场景,完全还原原始张量的形状与类型。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 22:35:40