Refactor: Correct and rename VideoContinuationGenerator start logic

- Rename 'beginning' to 'beginning_of_generation' for clarity
- Rename 'after overlap_frames' to 'after_overlap_frames'
- Swap underlying logic to match intuitive naming
- Set default to 'beginning_of_generation' for both parameters
- Update tooltips to accurately reflect the new behavior
This commit is contained in:
POM
2025-06-23 11:34:47 +02:00
parent 08cd413f0b
commit 12bdbfbc68
+23 -22
View File
@@ -822,7 +822,8 @@ class VideoContinuationGenerator:
"end_frame": ("IMAGE", {"tooltip": "Optional single frame to place at the end of the continuation video."}),
"control_images": ("IMAGE", {"tooltip": "Optional control images to fill the empty frames."}),
"inpaint_mask": ("MASK", {"tooltip": "Optional inpaint mask to use for the empty frames, overriding the default mask."}),
"continuation_mode": (["Generate new content", "Stitch to existing sequence"], {"default": "Generate new content", "tooltip": "Choose 'Generate new content' if your control images are only for the new section. Choose 'Stitch to existing sequence' if your control images represent the full, final timeline."}),
"when_to_start_control_frames": (["beginning_of_generation", "after_overlap_frames"], {"default": "beginning_of_generation", "tooltip": "Aligns control frames with the start of the output (frame 0). Overlap frames from the input video will take priority, so control frames will become visible starting after the overlap period."}),
"when_to_start_masks": (["beginning_of_generation", "after_overlap_frames"], {"default": "beginning_of_generation", "tooltip": "Aligns masks with the start of the output (frame 0). 'after_overlap_frames' starts masks after the overlap period."}),
},
}
@@ -832,7 +833,7 @@ class VideoContinuationGenerator:
CATEGORY = "Steerable-Motion"
DESCRIPTION = "Creates a continuation video by placing overlap frames from the end of input video at the start, with optional end frame."
def generate_continuation_video(self, input_video_frames, total_output_frames, overlap_frames, empty_frame_fill_level, end_frame=None, control_images=None, inpaint_mask=None, continuation_mode="Generate new content"):
def generate_continuation_video(self, input_video_frames, total_output_frames, overlap_frames, empty_frame_fill_level, end_frame=None, control_images=None, inpaint_mask=None, when_to_start_control_frames="beginning_of_generation", when_to_start_masks="beginning_of_generation"):
# 1. Validation and Setup
total_output_frames = int(total_output_frames)
if (total_output_frames - 1) % 4 != 0:
@@ -877,19 +878,10 @@ class VideoContinuationGenerator:
if num_middle_frames > 0:
if control_images is not None:
log.info(f"Using 'control_images' to fill the {num_middle_frames} middle frames with '{continuation_mode}' mode.")
log.info(f"Using 'control_images' to fill the {num_middle_frames} middle frames with '{when_to_start_control_frames}' mode.")
control_images_resized = common_upscale(control_images.movedim(-1, 1), frame_width, frame_height, "lanczos", "disabled").movedim(1, -1)
if continuation_mode == "Generate new content":
# Use control frames from the beginning of the sequence (C0, C1, C2...)
if control_images_resized.shape[0] < num_middle_frames:
log.warning(f"Provided 'control_images' have {control_images_resized.shape[0]} frames, less than needed ({num_middle_frames}). Padding with 'empty_frame_fill_level'.")
padding_needed = num_middle_frames - control_images_resized.shape[0]
padding = torch.ones((padding_needed, frame_height, frame_width, num_channels), device=device, dtype=dtype) * empty_frame_fill_level
middle_frames_part = torch.cat([control_images_resized, padding], dim=0)
else:
middle_frames_part = control_images_resized[:num_middle_frames].clone()
else: # "Stitch to existing sequence"
if when_to_start_control_frames == "beginning_of_generation":
# Skip the first overlap_frames control images to avoid duplication
duplicate_count = min(actual_overlap_frames, control_images_resized.shape[0])
available_after_dup = control_images_resized.shape[0] - duplicate_count
@@ -901,6 +893,15 @@ class VideoContinuationGenerator:
middle_frames_part = torch.cat([selected_control, padding], dim=0)
else:
middle_frames_part = control_images_resized[duplicate_count:duplicate_count + num_middle_frames].clone()
else: # "after_overlap_frames"
# Use control frames from the beginning of the sequence (C0, C1, C2...)
if control_images_resized.shape[0] < num_middle_frames:
log.warning(f"Provided 'control_images' have {control_images_resized.shape[0]} frames, less than needed ({num_middle_frames}). Padding with 'empty_frame_fill_level'.")
padding_needed = num_middle_frames - control_images_resized.shape[0]
padding = torch.ones((padding_needed, frame_height, frame_width, num_channels), device=device, dtype=dtype) * empty_frame_fill_level
middle_frames_part = torch.cat([control_images_resized, padding], dim=0)
else:
middle_frames_part = control_images_resized[:num_middle_frames].clone()
else:
log.info(f"No 'control_images', filling {num_middle_frames} middle frames with level {empty_frame_fill_level}.")
middle_frames_part = torch.ones((num_middle_frames, frame_height, frame_width, num_channels), device=device, dtype=dtype) * empty_frame_fill_level
@@ -911,14 +912,8 @@ class VideoContinuationGenerator:
# 6. Create Mask
continuation_frame_masks = torch.ones((total_output_frames, frame_height, frame_width), device=device, dtype=dtype)
# Apply mask logic based on continuation_mode parameter
if continuation_mode == "Generate new content":
# Set known frames (overlap and end) to 0.0, rest stay as 1.0 (inpaint)
if actual_overlap_frames > 0:
continuation_frame_masks[0:actual_overlap_frames] = 0.0
if num_end_frames > 0:
continuation_frame_masks[-num_end_frames:] = 0.0
else: # "Stitch to existing sequence"
# Apply mask logic based on when_to_start_masks parameter
if when_to_start_masks == "beginning_of_generation":
# Set known frames (overlap and end) to 0.0, but also set middle section based on control frame logic
if actual_overlap_frames > 0:
continuation_frame_masks[0:actual_overlap_frames] = 0.0
@@ -934,7 +929,13 @@ class VideoContinuationGenerator:
middle_start = actual_overlap_frames
middle_end = middle_start + num_middle_frames
continuation_frame_masks[middle_start:middle_end] = 0.0
else: # "after_overlap_frames"
# Set known frames (overlap and end) to 0.0, rest stay as 1.0 (inpaint)
if actual_overlap_frames > 0:
continuation_frame_masks[0:actual_overlap_frames] = 0.0
if num_end_frames > 0:
continuation_frame_masks[-num_end_frames:] = 0.0
# 7. Handle optional inpaint_mask override
if inpaint_mask is not None:
log.info("Processing provided 'inpaint_mask', which will override the automatically generated mask.")