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

如何重写TensorFlow的TFRecord解析函数以实现向量化批处理?

如何向量化TFRecord加载函数以支持批处理并解决形状未定义问题

我明白你现在的痛点——原来的单样本解析函数在大规模数据加载时效率不够,想改成批处理解析,但卡在了张量形状和批量解析的问题上。其实核心是要调整解析流程:从先单样本解析再批处理改成先批处理再批量解析,同时适配parse_example和批量张量解析的逻辑。

下面是修改后的完整可运行代码,我会在关键部分标注改动说明:

import tensorflow as tf
import os
import numpy as np
import tensorflow_datasets as tfds

AUTOTUNE = tf.data.experimental.AUTOTUNE
ds = tfds.load('mnist', shuffle_files=True, as_supervised=True)
ds['test'].cardinality()
ds_splits = ["train", "test"]

## Write features (这部分保持不变)
def _bytes_feature(value):
    """Returns a bytes_list from a string / byte."""
    if isinstance(value, type(tf.constant(0))): # if value ist tensor
        value = value.numpy() # get value of tensor
    return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))

def _float_feature(value):
    """Returns a floast_list from a float / double."""
    return tf.train.Feature(float_list=tf.train.FloatList(value=[value]))

def _int64_feature(value):
    """Returns an int64_list from a bool / enum / int / uint."""
    return tf.train.Feature(int64_list=tf.train.Int64List(value=[value]))

def serialize_array(array):
    array = tf.io.serialize_tensor(array)
    return array

for d in ds_splits:
    print("saving {}".format(d))
    subset = ds[d]
    filename = d+".tfrecords"
    writer = tf.io.TFRecordWriter(filename)
    count = 0
    for image, label in subset:
        data={
            'height': _int64_feature(28),
            'width': _int64_feature(28),
            'depth': _int64_feature(1),
            'label': _int64_feature(label),
            'image_raw':_bytes_feature(serialize_array(image))
        }
        out = tf.train.Example(features=tf.train.Features(feature=data))
        writer.write(out.SerializeToString())
        count +=1
    writer.close()
    print(count)

## Load features (核心修改部分)
def parse_tfr_batch(batch_elements):
    # 定义特征解析字典,和原来一致,但parse_example会处理批量数据
    parse_dict = {
        'height': tf.io.FixedLenFeature([], tf.int64),
        'width':tf.io.FixedLenFeature([], tf.int64),
        'label':tf.io.FixedLenFeature([], tf.int64),
        'depth':tf.io.FixedLenFeature([], tf.int64),
        'image_raw' : tf.io.FixedLenFeature([], tf.string)
    }
    # 用parse_example替代parse_single_example,处理批量元素
    example_messages = tf.io.parse_example(batch_elements, parse_dict)
    
    # 批量解析图像张量:用tf.map_fn对每个字符串调用parse_tensor
    # out_type指定输出类型,和序列化时一致
    images = tf.map_fn(
        lambda x: tf.io.parse_tensor(x, out_type=tf.uint8),
        example_messages['image_raw'],
        fn_output_signature=tf.TensorSpec(shape=(28,28,1), dtype=tf.uint8)
    )
    
    # 获取批次大小,构造动态形状
    batch_size = tf.shape(images)[0]
    # 这里因为所有图像尺寸固定,也可以直接reshape成[None,28,28,1]
    images = tf.reshape(images, shape=[batch_size, 28, 28, 1])
    
    labels = example_messages['label']
    return (images, labels)

def get_dataset(filename, set_type, batch_size=32):
    ignore_order = tf.data.Options()
    ignore_order.experimental_deterministic = False# disable native order, increase speed
    dataset = tf.data.TFRecordDataset(filename)
    dataset = dataset.with_options( ignore_order )
    
    # 关键改动:先batch再map,这样传给解析函数的是一批元素
    dataset = dataset.batch(batch_size)
    dataset = dataset.map(parse_tfr_batch, num_parallel_calls=AUTOTUNE)
    
    dataset = dataset.shuffle(2048, reshuffle_each_iteration=True)
    dataset = dataset.prefetch(buffer_size=AUTOTUNE)
    dataset = dataset.repeat() if set_type =='train' else dataset
    return dataset

BATCH_SIZE = 32
tfr_dataset = get_dataset('train.tfrecords', "train", batch_size=BATCH_SIZE)
for sample in tfr_dataset.take(1): #sanity check
    print("Image shape:", sample[0].shape)
    print("Label shape:", sample[1].shape)

## Training and evaluation (保持不变)
def get_cnn():
    model = tf.keras.Sequential([
        tf.keras.layers.Conv2D(kernel_size=3, filters=16, padding='same', activation='relu', input_shape=[28,28, 1]),
        tf.keras.layers.Conv2D(kernel_size=3, filters=32, padding='same', activation='relu'),
        tf.keras.layers.MaxPooling2D(pool_size=2),
        tf.keras.layers.Conv2D(kernel_size=3, filters=64, padding='same', activation='relu'),
        tf.keras.layers.MaxPooling2D(pool_size=2),
        tf.keras.layers.Conv2D(kernel_size=3, filters=128, padding='same', activation='relu'),
        tf.keras.layers.MaxPooling2D(pool_size=2),
        tf.keras.layers.Conv2D(kernel_size=3, filters=256, padding='same', activation='relu'),
        tf.keras.layers.GlobalAveragePooling2D(),
        tf.keras.layers.Dense(10,'softmax')
    ])
    optimizer = tf.keras.optimizers.RMSprop(lr=0.01)
    model.compile(loss='sparse_categorical_crossentropy', optimizer=optimizer, metrics=[tf.keras.metrics.SparseCategoricalAccuracy()])
    return model

model = get_cnn()
model.summary()
model.fit(tfr_dataset, steps_per_epoch=60000//BATCH_SIZE, epochs=2)

关键改动说明:

  • 解析函数从单样本改批量:parse_tfr_batch接收一批TFRecord元素,用tf.io.parse_example替代parse_single_example,返回的每个特征都带有批次维度(比如label形状是[batch_size])。
  • 批量解析图像张量:用tf.map_fn对批量的image_raw字符串逐个解析,通过fn_output_signature指定输出张量的形状和类型,让TensorFlow能确定张量的秩,避免模型报错。
  • 调整数据处理顺序:在get_dataset中先调用batch再调用map,这样解析函数直接处理批量数据,减少了单样本解析的开销,实现向量化加速。
  • 动态形状设置:利用tf.shape(images)[0]获取当前批次的大小,结合固定的图像尺寸构造输出形状,既适配不同批次大小,又保证了张量秩的确定性,解决了模型输入不兼容的问题。

运行这段代码后,你会发现模型能正常接收输入,不会再出现rank is undefined的错误,同时数据加载效率也会因为批处理解析得到提升。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 20:22:38