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

如何在TensorFlow序列化Tensor到TFRecord时保留数组形状?

解决TFRecord序列化时保留Tensor数组形状的问题

直接用numpy.bytes()序列化Tensor对应的numpy数组时,只会保存原始字节数据,丢失了数组形状的元信息,导致解码时无法直接恢复原数组结构。最优方案是在TFRecord的Feature字典中额外添加shape字段,存储数组的形状信息,具体实现如下:

1. 修改写入TFRecord的代码

调整特征生成函数以支持列表类型的形状数据,并在写入时新增形状字段:

import tensorflow as tf
from tqdm import tqdm

def _bytes_feature(value):
    # 从字符串/字节生成bytes_list类型的Feature
    if isinstance(value, type(tf.constant(0))):
        value = value.numpy()
    return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))

def _int64_feature(value):
    # 支持单个数值或数值列表,生成int64_list类型的Feature
    if not isinstance(value, (list, tuple)):
        value = [value]
    return tf.train.Feature(int64_list=tf.train.Int64List(value=value))

def write_tfrecords(data_list, output_file):
    """
    写入数据到TFRecord,同时保留图像数组形状
    """
    total_samples = 0
    with tf.io.TFRecordWriter(output_file) as writer:    
        for image, label in tqdm(data_list):
            img_np = image.numpy()
            data = {
                "image": _bytes_feature(img_np.tobytes()),
                "shape": _int64_feature(img_np.shape),  # 存储数组的形状元组
                "label": _int64_feature(label)
            }
            example = tf.train.Example(features=tf.train.Features(feature=data))
            writer.write(example.SerializeToString())
            total_samples += 1
    return total_samples

2. 对应的解码代码

解析TFRecord时读取形状信息,用它来恢复原数组结构:

def parse_tfrecord_example(example_proto):
    # 定义与写入时对应的Feature解析结构
    feature_desc = {
        "image": tf.io.FixedLenFeature([], tf.string),
        "shape": tf.io.VarLenFeature(tf.int64),  # 形状维度数不固定,用变长特征解析
        "label": tf.io.FixedLenFeature([], tf.int64)
    }
    parsed = tf.io.parse_single_example(example_proto, feature_desc)
    
    # 恢复图像数组:先解码字节为一维张量,再用原形状reshape
    # 注意:tf.io.decode_raw的dtype要与原数组一致,比如原数组是float32则改为tf.float32
    image_raw = tf.io.decode_raw(parsed["image"], tf.uint8)
    shape = tf.sparse.to_dense(parsed["shape"])  # 把稀疏张量转为密集张量
    image = tf.reshape(image_raw, shape)
    
    label = parsed["label"]
    return image, label

# 加载并解析TFRecord示例
dataset = tf.data.TFRecordDataset("your_output_path.tfrecords")
dataset = dataset.map(parse_tfrecord_example)

关键说明

  • 写入时通过img_np.shape获取数组形状,转为int64列表存入TFRecord,确保形状信息不丢失
  • 解码时用VarLenFeature解析形状(适配不同维度的数组,比如2D灰度图、3D彩图),再通过tf.reshape恢复原数组结构
  • 解码时tf.io.decode_raw的dtype必须与原数组的数据类型完全匹配,否则会出现数据错乱

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 21:06:30