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

如何在Keras中使用ImageDataGenerator对视频(4D张量)做数据增强?

如何用Keras的ImageDataGenerator增强视频帧序列

好问题!要把ImageDataGenerator的数据增强逻辑应用到短视频帧序列(形状为(num_frames, width, height, 3)的4D数组)上,核心要求是同一个视频的所有帧必须使用完全相同的增强参数——不然帧之间的时序连贯性会被破坏,模型训练效果会大打折扣。下面是具体的实现思路和代码:

核心思路

ImageDataGenerator默认是为单张图像(3D数组)生成随机增强变换的,我们可以利用它的get_random_transform和apply_transform方法,为每个视频单独生成一组固定的增强参数,再将这组参数统一应用到该视频的所有帧上。

具体实现步骤

1. 初始化ImageDataGenerator

和处理单张图像时一样,先定义你的增强规则:

import tensorflow as tf
import numpy as np

# 定义增强规则,和你之前的设置一致
aug = tf.keras.preprocessing.image.ImageDataGenerator(
    rotation_range=15,
    zoom_range=0.15
    # 可以添加其他增强,比如horizontal_flip=True等
)

2. 自定义视频数据生成器

写一个生成器函数,负责遍历视频数据集,为每个视频生成一致的增强帧:

def video_aug_generator(videos, labels, batch_size=32):
    num_videos = len(videos)
    while True:
        # 每次迭代前打乱数据,保证训练随机性
        indices = np.random.permutation(num_videos)
        for start_idx in range(0, num_videos, batch_size):
            # 截取当前batch的视频和标签
            end_idx = min(start_idx + batch_size, num_videos)
            batch_vids = videos[indices[start_idx:end_idx]]
            batch_lbls = labels[indices[start_idx:end_idx]]
            
            augmented_batch = []
            for vid in batch_vids:
                # 为当前视频生成一套固定的增强变换参数(基于第一帧的形状)
                transform = aug.get_random_transform(vid[0].shape)
                # 将这套参数应用到视频的每一帧
                augmented_frames = [aug.apply_transform(frame, transform) for frame in vid]
                augmented_batch.append(np.array(augmented_frames))
            
            # 返回增强后的视频batch和对应标签
            yield np.array(augmented_batch), np.array(batch_lbls)

3. 使用生成器训练模型

假设你的训练集X_train是形状为(num_videos, num_frames, width, height, 3)的5D数组,y_train是对应的标签数组,直接用自定义生成器训练即可:

train_gen = video_aug_generator(X_train, y_train, batch_size=32)
# 计算每个epoch需要执行的步数
steps_per_epoch = np.ceil(len(X_train) / 32).astype(int)

model.fit(
    train_gen,
    epochs=100,
    steps_per_epoch=steps_per_epoch
)

关键细节说明

  • 时序一致性:通过get_random_transform为每个视频生成唯一的变换参数,再用apply_transform统一作用于所有帧,确保视频的动作、场景连贯性不会被增强操作破坏;
  • 扩展性:如果需要添加其他增强方式(比如平移、翻转),直接在ImageDataGenerator的初始化参数里添加即可,get_random_transform会自动包含这些变换的参数;
  • 兼容性:这个生成器完全兼容Keras的fit方法,和你之前用fit_generator的逻辑一致(注意Keras 2.1.0+推荐直接用fit替代fit_generator)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:25:17