diff --git a/nodes.py b/nodes.py index f7a4278..16d39d6 100644 --- a/nodes.py +++ b/nodes.py @@ -302,6 +302,7 @@ class DynamiCrafterBatchInterpolation: ], { "default": 'auto' }), + "cut_near_keyframes": ("INT", {"default": 0, "min": 0, "max": 5, "step": 1}), }, } @@ -310,7 +311,7 @@ class DynamiCrafterBatchInterpolation: FUNCTION = "process" CATEGORY = "DynamiCrafterWrapper" - def process(self, model, images, prompt, cfg, steps, eta, seed, fs, keep_model_loaded, frames, vae_dtype): + def process(self, model, images, prompt, cfg, steps, eta, seed, fs, keep_model_loaded, frames, vae_dtype, cut_near_keyframes): assert images.shape[0] > 1, "DynamiCrafterBatchInterpolation needs at least 2 images" device = mm.get_torch_device() mm.unload_all_models() @@ -453,6 +454,17 @@ class DynamiCrafterBatchInterpolation: out_video = F.interpolate(out_video.permute(0, 3, 1, 2), size=(orig_H, orig_W), mode="bicubic") out_video = video.permute(0, 2, 3, 1) + # should we trim middle keyframes? + if cut_near_keyframes > 0: + already_deleted = 0 + for i in range(len(images) - 2): + old_size = out_video.shape[0] + keyframe_index = (i + 1) * frames - already_deleted + start_index = keyframe_index - (cut_near_keyframes // 2) + end_index = start_index + cut_near_keyframes + out_video = torch.cat([out_video[:start_index], out_video[end_index:]], dim=0) + already_deleted += old_size - out_video.shape[0] + last_image = out_video[-1].unsqueeze(0) return (out_video, last_image)