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

