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:
+23
-22
@@ -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.")
|
||||
|
||||
Reference in New Issue
Block a user