支持设置采样数量

This commit is contained in:
shadowcz007
2024-09-12 11:16:22 +08:00
parent e95b3813cc
commit cee54f336e
+64 -19
View File
@@ -3,22 +3,20 @@ import folder_paths
import numpy as np
import torch
from PIL import Image
import comfy.utils
from PIL import Image, ImageOps
from PIL.PngImagePlugin import PngInfo
# import comfy.utils
from PIL import Image
# from PIL.PngImagePlugin import PngInfo
import cv2
from scenedetect.video_manager import VideoManager
from scenedetect.scene_manager import SceneManager
from scenedetect.detectors import AdaptiveDetector
import os
import random
import string
class AnyType(str):
"""A special class that is always equal in not equal comparisons."""
@@ -87,8 +85,42 @@ def detect_scenes(video_path, min_scene_len=15, adaptive_threshold=3.0,callback=
return scenes
# 提取每个场景中的中间帧
def split_video_by_scenes(video_path, scenes, output_path):
# 采样逻辑
def calculate_sample_range(start_frame, middle_frame, end_frame, number_of_sample_frames):
half_samples = number_of_sample_frames // 2
# 初始化采样帧列表
samples = [middle_frame]
# 计算间隔
interval_before = (middle_frame - start_frame) // half_samples
interval_after = (end_frame - middle_frame) // half_samples
# 添加中间帧前的采样帧
for i in range(1, half_samples + 1):
sample_before = middle_frame - i * interval_before
if sample_before >= start_frame:
samples.insert(0, sample_before)
# 添加中间帧后的采样帧
for i in range(1, half_samples + 1):
sample_after = middle_frame + i * interval_after
if sample_after <= end_frame:
samples.append(sample_after)
# 如果采样帧数是偶数,则需要移除最靠近边界的一个帧
if number_of_sample_frames % 2 == 0:
if len(samples) > number_of_sample_frames:
if abs(samples[0] - start_frame) < abs(samples[-1] - end_frame):
samples.pop(0)
else:
samples.pop()
return samples
def split_video_by_scenes(video_path, scenes, output_path, number_of_sample_frames=1):
# Load the video file
video = cv2.VideoCapture(video_path)
@@ -113,6 +145,13 @@ def split_video_by_scenes(video_path, scenes, output_path):
# Calculate the middle frame
middle_frame = (start_frame + end_frame) // 2
sample_frames=[]
# Calculate the range of frames to sample
# sample_range = range(max(start_frame, middle_frame - number_of_sample_frames // 2),
# min(end_frame, middle_frame + number_of_sample_frames // 2 + 1))
sample_range=calculate_sample_range(start_frame, middle_frame, end_frame, number_of_sample_frames)
# Set the video file's current frame to the start frame
video.set(cv2.CAP_PROP_POS_FRAMES, start_frame)
@@ -128,22 +167,24 @@ def split_video_by_scenes(video_path, scenes, output_path):
break
writer.write(frame)
# If this is the middle frame, save it to keyframes
if frame_num == middle_frame:
# If this frame is in the sample range, save it to keyframes
if frame_num in sample_range:
# Convert the frame to RGB (OpenCV uses BGR by default)
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
# Convert the frame to a PIL image
pil_image = Image.fromarray(frame_rgb)
keyframe_info = {
sample_frames.append(pil2tensor(pil_image))
keyframe_info = {
'start_frame': start_frame,
'end_frame': end_frame,
'middle_frame': pil2tensor(pil_image),
'video_path':output_path1
'sample_frames': sample_frames,
'video_path': output_path1
}
keyframes.append(keyframe_info)
keyframes.append(keyframe_info)
# Release the VideoWriter object
writer.release()
@@ -176,7 +217,7 @@ class SceneInfoNode:
},}
RETURN_TYPES = ('IMAGE','INT','INT','SCENE_VIDEO',)
RETURN_NAMES = ("middle_frame","start_frame","end_frame","scene_video",)
RETURN_NAMES = ("sample_frames","start_frame","end_frame","scene_video",)
# OUTPUT_IS_LIST = (False,)
FUNCTION = "run"
@@ -191,7 +232,8 @@ class SceneInfoNode:
video_paths=[]
for i in range(len(scenes)):
s=scenes[i]
images.append(s['middle_frame'])
for sf in s['sample_frames']:
images.append(sf)
start_frames.append(s['start_frame'])
end_frames.append(s['end_frame'])
video_paths.append(s['video_path'])
@@ -199,7 +241,9 @@ class SceneInfoNode:
return (images,start_frames,end_frames,video_paths,)
else:
s=scenes[index]
return (s['middle_frame'],s['start_frame'],s['end_frame'],s['video_path'],)
images=s['sample_frames']
images = torch.cat(images, dim=0)
return (images,s['start_frame'],s['end_frame'],s['video_path'],)
# 分割视频
@@ -218,6 +262,7 @@ class ScenedetectNode_:
"video": (sorted(files), {"video_upload": True}),
"min_scene_len": ("INT", {"default": 10, "min": 1, "step": 1}),
"adaptive_threshold": ("FLOAT", {"default": 2.5, "min": 0, "step": 0.1}),
"number_of_sample_frames": ("INT", {"default": 1, "min": 1, "step": 1}), # 抽取的帧数,默认是1帧,中间帧
},}
RETURN_TYPES = ("SCENE_VIDEO","SCENE_", "INT",)
@@ -227,7 +272,7 @@ class ScenedetectNode_:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Video"
def run(self, video, min_scene_len,adaptive_threshold):
def run(self, video, min_scene_len,adaptive_threshold,number_of_sample_frames):
video_path = folder_paths.get_annotated_filepath(video)
# Example usage:
scenes = detect_scenes(video_path, min_scene_len=min_scene_len, adaptive_threshold=adaptive_threshold)
@@ -243,7 +288,7 @@ class ScenedetectNode_:
folder_path = create_folder(tp,name_without_extension)
# print("New folder created:", folder_path)
vs_files,keyframes=split_video_by_scenes(video_path,scenes,folder_path)
vs_files,keyframes=split_video_by_scenes(video_path,scenes,folder_path,number_of_sample_frames)
# print("New folder created:", vs_files)
return (vs_files,keyframes,len(scenes),)