如何加速tf.dataset的创建与训练?超10万帧视频数据集优化咨询
优化10万+视频帧数据集的创建速度
我正在为训练任务创建带标签的视频帧数据集,需要用特定数字标注帧中圆点的位置。视频总长约1小时、帧率25fps,总帧数超过10万,目前创建数据集耗时已经超过一天。我用tf.keras.image_dataset_from_directory从目录读取文件,再通过numpy数组分配标签,想问问有没有更快的数据集创建方式,或者并行合并tf.data.Datasets的方法。
我的代码如下(使用VSCode的Jupyter Notebook,Python3.9.7,TensorFlow 2.0):
import numpy as np import tensorflow as tf def print_dataset(data_set:tf.data.Dataset): iterator = data_set.as_numpy_iterator() # create a numpy array labelled images = np.empty((0, 108, 192, 1)) labels = np.empty((0)) for element in iterator: images = np.append(images, element[0], axis=0) labels = np.append(labels, element[1], axis=0) print(np.shape(images)) print(np.shape(labels)) plt.imshow(images[0]) print(labels[0]) return images, labels # creating the frame dataset from directory image = tf.keras.utils.image_dataset_from_directory( 'frames', labels=labels, label_mode='int', image_size=(108,192), color_mode='grayscale', batch_size=1) # frames directory contains the frames used for training. print(image) plt.show(image) #generating the labelled dataset and getting the img as well as labels img, labels= print_dataset(image)
核心优化方案
1. 彻底抛弃Numpy数组拼接的低效操作
你当前代码里的print_dataset函数是性能黑洞——np.append每次都会重新分配内存并复制全部已有数据,10万帧的场景下会导致指数级的时间消耗。完全不需要把tf.data.Dataset转换成Numpy数组,直接用Dataset对象进行后续训练即可。
2. 提升image_dataset_from_directory的读取效率
- 增大
batch_size:把当前的batch_size=1改成64/128(根据显存容量调整),批量读取能大幅降低磁盘IO的开销。 - 开启并行预处理:添加
num_parallel_calls=tf.data.AUTOTUNE参数,让TensorFlow并行加载和处理图像。 - 预取数据:在数据集末尾加上
.prefetch(tf.data.AUTOTUNE),让模型训练的同时提前加载下一批数据,消除数据等待间隙。
3. 并行生成标签
如果标签可以通过文件名、文件路径等规则生成,直接在Dataset流水线里并行处理,不用先读完全部数据再分配标签。
优化后的代码示例
import tensorflow as tf import matplotlib.pyplot as plt import os # 根据你的实际规则定义标签生成函数 def extract_label(file_path): # 示例:假设文件名格式为"frame_xxx_label.jpg",从文件名提取标签 filename = tf.strings.split(file_path, os.sep)[-1] label_str = tf.strings.split(filename, "_")[-1] label_str = tf.strings.split(label_str, ".")[0] # 去掉文件后缀 return tf.strings.to_number(label_str, out_type=tf.int32) def build_efficient_dataset(): # 1. 获取所有图像文件路径 file_paths = tf.data.Dataset.list_files('frames/**/*.jpg', shuffle=False) # 匹配所有子目录下的图像 # 2. 并行加载图像并生成标签 dataset = file_paths.map( lambda path: ( tf.keras.utils.load_img(path, target_size=(108,192), color_mode='grayscale'), extract_label(path) ), num_parallel_calls=tf.data.AUTOTUNE ) # 3. 转换图像为张量并归一化 dataset = dataset.map( lambda img, label: (tf.convert_to_tensor(img, dtype=tf.float32)/255.0, label), num_parallel_calls=tf.data.AUTOTUNE ) # 4. 批量处理+预取 dataset = dataset.batch(64).prefetch(tf.data.AUTOTUNE) return dataset # 使用优化后的数据集 train_dataset = build_efficient_dataset() # 验证数据(取一个batch查看) for batch_imgs, batch_labels in train_dataset.take(1): print(f"Batch图像形状: {batch_imgs.shape}") print(f"Batch标签形状: {batch_labels.shape}") plt.imshow(batch_imgs[0].numpy().squeeze(), cmap='gray') plt.title(f"标签: {batch_labels[0].numpy()}") plt.show()
额外提速建议
- 转换成TFRecord格式:将所有图像和标签打包成TFRecord文件,一次预处理后,后续读取速度会提升数倍,适合大规模数据集。
- 缓存数据集:在流水线中添加
.cache(),将数据集缓存到内存或磁盘,避免重复加载和预处理。 - 优化图像格式:如果当前用的是PNG等无损格式,转成JPG可以减小文件体积,加快读取速度。
内容的提问来源于stack exchange,提问作者user9228288
相关产品推荐
相关产品推荐

