如何使用视频数据训练模型?原图像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
相关产品推荐
相关产品推荐

