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

如何修改代码实现每个视频处理完成后生成独立H5特征文件

代码修改方案

核心修改逻辑

原来的代码初始化时就创建了全局H5文件,预生成所有视频的存储分组,所有视频数据都写入同一个文件。我们要调整为:

  • 初始化阶段仅加载视频列表、初始化特征提取模型,不提前创建全局H5文件
  • 每处理完一个视频,单独生成以该视频名命名的H5文件,直接存储该视频的所有特征数据

1. 修改Generate_Dataset.py

完整修改后的代码如下:

import os
from networks.CNN import ResNet
from utils.KTS.cpd_auto import cpd_auto
from tqdm import tqdm
import math
import cv2
import numpy as np
import h5py

class Generate_Dataset:
    def __init__(self, video_path, save_dir):
        self.resnet = ResNet()
        self.dataset = {}
        self.video_list = []
        self.video_path = ''
        self.save_dir = save_dir # 保存目录,不再是单个文件路径

        self._set_video_list(video_path)

    def _set_video_list(self, video_path):
        if os.path.isdir(video_path):
            self.video_path = video_path
            fileExt = (".mp4",".avi")
            self.video_list = [_ for _ in os.listdir(video_path) if _.endswith(fileExt)]
            self.video_list.sort()
        else:
            self.video_path = ''
            self.video_list.append(video_path)
        # 移除原来预创建全局H5分组的代码

    def _extract_feature(self, frame):
        frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
        frame = cv2.resize(frame, (224, 224))
        res_pool5 = self.resnet(frame)
        frame_feat = res_pool5.cpu().data.numpy().flatten()
        return frame_feat

    def _get_change_points(self, video_feat, n_frame, fps):
        n = n_frame / fps
        m = int(math.ceil(n/2.0))
        K = np.dot(video_feat, video_feat.T)
        change_points, _ = cpd_auto(K, m, 1)
        change_points = np.concatenate(([0], change_points, [n_frame-1]))

        temp_change_points = []
        for idx in range(len(change_points)-1):
            segment = [change_points[idx], change_points[idx+1]-1]
            if idx == len(change_points)-2:
                segment = [change_points[idx], change_points[idx+1]]
            temp_change_points.append(segment)
        change_points = np.array(list(temp_change_points))

        arr = change_points
        list1 = arr.tolist()
        list2 = list1[-1].pop(1)
        cps_m = math.floor(arr[-1][1]/15)
        list1[-1].append(cps_m)
        arr = np.asarray(list1)
        arrmul = arr * 15
        median_frame = []
        for x in arrmul:
          med = np.mean(x)
          int_array = med.astype(int)
          median_frame.append(int_array)
        return arrmul

    def generate_dataset(self):
        print('[INFO] CNN processing')
        for video_idx, video_filename in enumerate(self.video_list):
            video_path = video_filename
            if os.path.isdir(self.video_path):
                video_path = os.path.join(self.video_path, video_filename)
            video_basename = os.path.basename(video_path).split('.')[0]
            # 处理视频提取特征
            video_capture = cv2.VideoCapture(video_path)
            fps = video_capture.get(cv2.CAP_PROP_FPS)
            n_frames = int(video_capture.get(cv2.CAP_PROP_FRAME_COUNT))
            picks = []
            video_feat = None
            video_feat_for_train = None
            for frame_idx in tqdm(range(n_frames-1)):
                success, frame = video_capture.read()
                if frame_idx % 15 == 0:
                    if success:
                        frame_feat = self._extract_feature(frame)                    
                        picks.append(frame_idx)
                        if video_feat_for_train is None:
                            video_feat_for_train = frame_feat
                        else:
                            video_feat_for_train = np.vstack((video_feat_for_train, frame_feat))
                        if video_feat is None:
                            video_feat = frame_feat
                        else:
                            video_feat = np.vstack((video_feat, frame_feat))
                    else:
                        break
            video_capture.release()
            arrmul = self._get_change_points(video_feat, n_frames, fps)
            
            # 单独生成当前视频的H5文件
            h5_save_path = os.path.join(self.save_dir, f"{video_basename}.h5")
            with h5py.File(h5_save_path, 'w') as h5_file:
                # 直接在H5根路径存储当前视频数据,不需要嵌套video_x分组,也可以根据需要保留分组
                h5_file['features'] = list(video_feat_for_train)
                h5_file['picks'] = np.array(list(picks))
                h5_file['n_frames'] = n_frames
                h5_file['fps'] = fps
                h5_file['video_name'] = video_basename
                h5_file['change_points'] = arrmul
            print(f"[INFO] 视频{video_filename}处理完成,特征已保存到{h5_save_path}")

2. 修改Create_data.py

完整修改后的代码如下:

import argparse
import os
from utils.generate_dataset import Generate_Dataset

# 修复原代码的引号缺失语法错误
parser = argparse.ArgumentParser("Welcome you to fraction")
# Dataset options
parser.add_argument('--input', '--split', type=str, help="输入视频文件或视频文件夹路径")
parser.add_argument('--output', type=str, default='./h5_output', help="H5文件输出目录路径")

args = parser.parse_args()
if __name__ == "__main__":
    # 自动创建输出目录
    os.makedirs(args.output, exist_ok=True)
    gen = Generate_Dataset(args.input, args.output)
    gen.generate_dataset()

使用说明

运行命令示例:

python Create_data.py --input 你的视频文件夹路径 --output 你要保存H5的文件夹路径

运行后会在输出目录下为每个视频生成对应的独立H5文件,文件名和原视频名一致。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 18:54:04