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

如何加速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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 15:10:24