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

如何使用视频数据训练模型?原图像Dataset适配方案失效求助

视频数据集加载与训练修正方案

现有问题分析

当前的VideoDataset存在几个核心问题:

  • 单个实例仅对应一个视频,无法批量加载client和imposter两类视频列表
  • 初始化时提前打开cv2.VideoCapture,在多进程DataLoader中会因无法序列化导致报错
  • 直接按帧索引返回样本,不同视频帧数不一致,会导致批处理时张量形状不匹配

修正后的VideoDataset实现

以下是适配你的数据格式(client/imposter文本列表)的视频数据集类,同时解决上述问题:

import cv2
import torch
from torch.utils.data import Dataset
from PIL import Image
import random

class VideoDataset(Dataset):
    def __init__(self, client_file: str, imposter_file: str, transforms=None, sample_frames=16):
        # 读取两类视频路径列表
        with open(client_file, "r") as f:
            self.client_videos = f.read().splitlines()
        with open(imposter_file, "r") as f:
            self.imposter_videos = f.read().splitlines()
        
        # 合并视频路径与对应标签
        self.video_paths = self.client_videos + self.imposter_videos
        self.labels = torch.cat((torch.ones(len(self.client_videos)), 
                                torch.zeros(len(self.imposter_videos))))
        
        self.transforms = transforms
        self.sample_frames = sample_frames  # 每个视频固定采样的帧数

    def __len__(self):
        return len(self.video_paths)

    def _sample_frames(self, video_cap):
        total_frames = int(video_cap.get(cv2.CAP_PROP_FRAME_COUNT))
        # 生成要采样的帧索引(均匀采样或随机采样)
        if total_frames <= self.sample_frames:
            # 帧数不足时重复采样
            frame_indices = list(range(total_frames)) + random.choices(range(total_frames), k=self.sample_frames - total_frames)
        else:
            # 均匀采样
            step = total_frames // self.sample_frames
            frame_indices = [i * step for i in range(self.sample_frames)]
            # 也可以用随机采样:frame_indices = random.sample(range(total_frames), self.sample_frames)
        
        frames = []
        for idx in frame_indices:
            video_cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
            success, frame = video_cap.read()
            if success:
                frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
                frame = Image.fromarray(frame)  # 转为PIL格式适配大部分transforms
                frames.append(frame)
        return frames

    def __getitem__(self, idx):
        video_path = self.video_paths[idx]
        label = self.labels[idx]
        
        # 延迟打开视频,避免多进程序列化问题
        video_cap = cv2.VideoCapture(video_path)
        if not video_cap.isOpened():
            raise ValueError(f"无法打开视频文件: {video_path}")
        
        # 采样固定数量的帧
        frames = self._sample_frames(video_cap)
        video_cap.release()  # 及时释放资源
        
        # 对每帧应用transforms
        if self.transforms:
            frames = [self.transforms(frame) for frame in frames]
        
        # 将帧堆叠成张量 (T, C, H, W),T为采样帧数
        video_tensor = torch.stack(frames)
        return video_tensor, label

数据加载器使用示例

和原图像数据集的使用方式保持一致:

# 假设preprocess是针对单帧的transforms(如Resize、ToTensor、Normalize等)
train_dataset = VideoDataset(
    client_file="/kaggle/input/dfdcdfdc/client_train_raw.txt",
    imposter_file="/kaggle/input/dfdcdfdc/imposter_train_raw.txt",
    transforms=preprocess,
    sample_frames=16
)
val_dataset = VideoDataset(
    client_file="/kaggle/input/dfdcdfdc/client_test_raw.txt",
    imposter_file="/kaggle/input/dfdcdfdc/imposter_test_raw.txt",
    transforms=preprocess,
    sample_frames=16
)

# 创建数据加载器,注意num_workers建议设为0或根据环境调整,避免cv2多进程问题
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=0)
val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=4, shuffle=False, num_workers=0)

关键注意事项

  • 多进程兼容性:cv2.VideoCapture在多进程环境下容易出问题,建议num_workers设为0,或使用multiprocessing.get_context('spawn')初始化DataLoader
  • Transforms适配:确保你的preprocess是针对单帧(PIL图像或numpy数组)的处理逻辑,不要包含针对视频序列的操作
  • 帧采样策略:可以根据任务需求调整_sample_frames方法,比如随机采样、取关键帧等,保证输入模型的序列长度一致

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 02:24:56