From 447518842e128862928ddcee8e906f462f375828 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 20 Mar 2024 00:14:03 +0200 Subject: [PATCH] Squashed commit of the following: commit 3fdd219fd73a5b0297e2751d423e73c03103e394 Merge: 1f2662a 2968761 Author: Phr00t Date: Sun Mar 17 22:17:39 2024 -0400 Merge branch 'main' into bugfix commit 1f2662a83b1c9dbdf7a487514ea00c6df369c09d Author: Phr00t Date: Sun Mar 17 20:14:51 2024 -0400 batch resizing fix commit 6ae23034e15d62978726eaa32e412380bbdf27f5 Author: Phr00t Date: Sun Mar 17 20:14:00 2024 -0400 cut out some frames around keyframes option to get rid of pauses --- nodes.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/nodes.py b/nodes.py index 2ed2866..5bbfb7c 100644 --- a/nodes.py +++ b/nodes.py @@ -310,6 +310,7 @@ class DynamiCrafterBatchInterpolation: ], { "default": 'auto' }), + "cut_near_keyframes": ("INT", {"default": 0, "min": 0, "max": 5, "step": 1}), }, } @@ -318,7 +319,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() @@ -465,6 +466,17 @@ class DynamiCrafterBatchInterpolation: if out_video.shape[1] != final_H or out_video.shape[2] != final_W: out_video = F.interpolate(out_video.permute(0, 3, 1, 2), size=(final_H, final_W), mode="bicubic").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)