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

