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

预处理HMDB51动作识别数据集时出现内存错误如何解决

内存错误根因

触发OOM的核心原因是一次性将所有视频提取的帧全部加载到内存。按当前参数计算:单个视频序列大小为20帧 * 64*64分辨率 *3通道 *4字节(float32)≈983KB,HMDB51全量51类共约6800个有效视频,总内存占用超过6.5GB,叠加Python列表、numpy数组的额外开销,很容易打满内存触发报错。

视频数据单样本维度远高于图像,Keras没有内置开箱即用的video_dataset_from_directory接口,必须自定义懒加载逻辑:仅在需要读取批次数据时才执行帧提取,不要提前把所有数据预处理完存入内存。

实现方案

1. 预先生成视频路径与标签映射

这一步仅存储文件路径和分类标签,内存占用可以忽略,同时提前过滤掉帧数不足的无效视频:

import os
import numpy as np
import cv2
import tensorflow as tf
from tensorflow import keras

# 原有常量定义保持不变
IMAGE_HEIGHT , IMAGE_WIDTH = 64, 64
SEQUENCE_LENGTH = 20
DATASET_DIR = r"\HMDB51"
CLASSES_LIST = ["brush_hair", "cartwheel", "catch", "chew", "clap", "climb", "climb_stairs", "dive",
            "draw_sword", "dribble", "drink", "eat", "fall_floor", "fencing", "flic_flac", "golf",
            "handstand", "hit", "hug", "jump", "kick", "kick_ball", "kiss", "laugh", 
            "pick", "pour", "pullup", "punch", "push", "pushup", "ride_bike", "ride_horse", 
            "run", "shake_hands", "shoot_ball", "shoot_bow", "shoot_gun", "sit", "situp", "smile", 
            "smoke", "somersault", "stand","swing_baseball", "sword", "sword_exercise", "talk", "throw", "turn", 
            "walk", "wave"]

def get_video_paths_labels():
    video_paths = []
    labels = []
    for class_index, class_name in enumerate(CLASSES_LIST):
        class_dir = os.path.join(DATASET_DIR, class_name)
        for file_name in os.listdir(class_dir):
            video_path = os.path.join(class_dir, file_name)
            # 提前校验视频帧数,避免训练时反复判断
            cap = cv2.VideoCapture(video_path)
            frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
            cap.release()
            if frame_count >= SEQUENCE_LENGTH:
                video_paths.append(video_path)
                labels.append(class_index)
    return np.array(video_paths), np.array(labels)

video_paths, labels = get_video_paths_labels()

2. 封装懒加载帧提取逻辑

将原有帧提取函数改造为适配tf.data流水线的形式,仅在数据被调用时执行解码、缩放、归一化操作:

def load_video_frames(video_path):
    frames_list = []
    # tf传入的路径为bytes类型,需要先解码
    video_reader = cv2.VideoCapture(video_path.decode("utf-8"))
    video_frames_count = int(video_reader.get(cv2.CAP_PROP_FRAME_COUNT))
    skip_frames_window = max(int(video_frames_count/SEQUENCE_LENGTH), 1)
    
    for frame_counter in range(SEQUENCE_LENGTH):
        video_reader.set(cv2.CAP_PROP_POS_FRAMES, frame_counter * skip_frames_window)
        success, frame = video_reader.read()
        if not success:
            break
        resized_frame = cv2.resize(frame, (IMAGE_HEIGHT, IMAGE_WIDTH))
        normalized_frame = resized_frame / 255.0
        frames_list.append(normalized_frame)
    
    video_reader.release()
    return np.array(frames_list, dtype=np.float32)

# 适配TensorFlow计算图的包装层
def tf_load_video(video_path, label):
    frames = tf.py_function(load_video_frames, inp=[video_path], Tout=tf.float32)
    # 手动指定张量形状,避免TensorFlow形状推导失败
    frames.set_shape((SEQUENCE_LENGTH, IMAGE_HEIGHT, IMAGE_WIDTH, 3))
    return frames, label

3. 构建分批次加载的数据集流水线

最终生成的数据集和keras.utils.image_dataset_from_directory返回的对象用法完全一致,支持多线程并行预处理、预加载等优化:

# 配置参数
BATCH_SIZE = 8
SHUFFLE_BUFFER_SIZE = 500 # 打乱缓冲区大小,无需设为总样本数避免占用内存

# 构建数据集
dataset = tf.data.Dataset.from_tensor_slices((video_paths, labels))
dataset = dataset.shuffle(SHUFFLE_BUFFER_SIZE)
# 多线程并行加载视频帧,不阻塞GPU训练
dataset = dataset.map(tf_load_video, num_parallel_calls=tf.data.AUTOTUNE)
dataset = dataset.batch(BATCH_SIZE)
# 预加载下一批次数据,训练和数据加载并行
dataset = dataset.prefetch(tf.data.AUTOTUNE)

# 划分训练集、验证集
train_size = int(0.8 * len(video_paths))
train_ds = dataset.take(train_size)
val_ds = dataset.skip(train_size)

4. 训练阶段直接传入数据集

原有模型结构无需修改,训练时直接传入数据集对象即可:

model.fit(
    train_ds,
    validation_data=val_ds,
    epochs=30
)
优化建议
  • 如果训练时CPU解码视频速度跟不上GPU,可以提前将所有视频提取的帧存为单个npy文件到磁盘,训练时直接加载npy文件,速度比实时解码快3~5倍
  • 批大小根据显存调整,64*64分辨率的输入在8G显存的GPU上可设置BATCH_SIZE=8~16
  • 内存富余时可适当调大SHUFFLE_BUFFER_SIZE,数据打乱效果更好

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 01:15:35