diff --git a/SteerableMotion.py b/SteerableMotion.py index 5604215..c0b6810 100644 --- a/SteerableMotion.py +++ b/SteerableMotion.py @@ -822,8 +822,7 @@ 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."}), - "when_to_start_control_frames": (["beginning", "after overlap_frames"], {"default": "after overlap_frames", "tooltip": "If at beginning, control frames won't actually be active until after the context, but after that they'll continue on from after the overlap_frames number. If after overlap_frames, the first control frame will be placed after the context is over."}), - "when_to_start_masks": (["beginning", "after overlap_frames"], {"default": "after overlap_frames", "tooltip": "Controls when mask generation begins. If at beginning, masks start from frame 0. If after overlap_frames, masks start after the overlap period to match control frame timing."}), + "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."}), }, } @@ -833,7 +832,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, when_to_start_control_frames="after overlap_frames", when_to_start_masks="after overlap_frames"): + 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"): # 1. Validation and Setup total_output_frames = int(total_output_frames) if (total_output_frames - 1) % 4 != 0: @@ -878,11 +877,11 @@ 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 '{when_to_start_control_frames}' mode.") + log.info(f"Using 'control_images' to fill the {num_middle_frames} middle frames with '{continuation_mode}' mode.") control_images_resized = common_upscale(control_images.movedim(-1, 1), frame_width, frame_height, "lanczos", "disabled").movedim(1, -1) - if when_to_start_control_frames == "beginning": - # Start from the beginning of control_images, regardless of overlap + 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] @@ -890,12 +889,12 @@ class VideoContinuationGenerator: middle_frames_part = torch.cat([control_images_resized, padding], dim=0) else: middle_frames_part = control_images_resized[:num_middle_frames].clone() - else: # "after overlap_frames" - # Skip potential duplicate frames that overlap with the start section + else: # "Stitch to existing sequence" + # 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 if available_after_dup < num_middle_frames: - log.info(f"After removing {duplicate_count} overlapping frames, only {available_after_dup} control frames remain; padding {num_middle_frames - available_after_dup} frames with 'empty_frame_fill_level'.") + log.info(f"After skipping {duplicate_count} control frames, only {available_after_dup} remain; padding {num_middle_frames - available_after_dup} frames with 'empty_frame_fill_level'.") selected_control = control_images_resized[duplicate_count:] padding_needed = num_middle_frames - selected_control.shape[0] padding = torch.ones((padding_needed, frame_height, frame_width, num_channels), device=device, dtype=dtype) * empty_frame_fill_level @@ -912,14 +911,14 @@ 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 when_to_start_masks parameter - if when_to_start_masks == "beginning": + # 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: # "after overlap_frames" + else: # "Stitch to existing sequence" # 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 diff --git a/video_continuation_comprehensive_simulation.png b/video_continuation_comprehensive_simulation.png new file mode 100644 index 0000000..bc492d3 Binary files /dev/null and b/video_continuation_comprehensive_simulation.png differ diff --git a/video_continuation_simulation.py b/video_continuation_simulation.py new file mode 100644 index 0000000..caadfb8 --- /dev/null +++ b/video_continuation_simulation.py @@ -0,0 +1,265 @@ +import matplotlib.pyplot as plt +import numpy as np + +def simulate_video_continuation_comprehensive(input_frames, control_frames, total_output_frames, overlap_frames, + when_to_start_control_frames, when_to_start_masks, end_frame=None): + """ + Comprehensive simulation of VideoContinuationGenerator logic + """ + print(f"\n=== Simulation: control='{when_to_start_control_frames}', masks='{when_to_start_masks}' ===") + + # Step 1: Calculate actual overlap frames + actual_overlap_frames = min(overlap_frames, len(input_frames), total_output_frames) + + # Step 2: Prepare start frames (from overlap) + overlap_start_idx = len(input_frames) - actual_overlap_frames + start_frames = input_frames[overlap_start_idx:overlap_start_idx + actual_overlap_frames] + + # Step 3: Prepare end frame + num_end_frames = 1 if end_frame is not None and total_output_frames > actual_overlap_frames else 0 + end_frames = [end_frame] if num_end_frames > 0 else [] + + # Step 4: Calculate middle frames needed + num_middle_frames = total_output_frames - actual_overlap_frames - num_end_frames + + # Step 5: Fill middle frames based on control frame mode + middle_frames = [] + control_frame_info = "" + + if num_middle_frames > 0 and control_frames: + if when_to_start_control_frames == "beginning": + # Use control frames from the beginning of the sequence (C0, C1, C2...) + if len(control_frames) < num_middle_frames: + middle_frames = control_frames + ['EMPTY'] * (num_middle_frames - len(control_frames)) + control_frame_info = f"Using first {len(control_frames)} control frames + {num_middle_frames - len(control_frames)} empty" + else: + middle_frames = control_frames[:num_middle_frames] + control_frame_info = f"Using first {num_middle_frames} control frames (C0-C{num_middle_frames-1})" + else: # "after overlap_frames" + # Skip the first overlap_frames control images to avoid duplication + duplicate_count = min(actual_overlap_frames, len(control_frames)) + available_after_dup = len(control_frames) - duplicate_count + + if available_after_dup < num_middle_frames: + selected_control = control_frames[duplicate_count:] + padding_needed = num_middle_frames - len(selected_control) + middle_frames = selected_control + ['EMPTY'] * padding_needed + control_frame_info = f"Skipped first {duplicate_count} control frames, using C{duplicate_count}-C{duplicate_count + len(selected_control) - 1} + {padding_needed} empty" + else: + middle_frames = control_frames[duplicate_count:duplicate_count + num_middle_frames] + control_frame_info = f"Skipped first {duplicate_count} control frames, using C{duplicate_count}-C{duplicate_count + num_middle_frames - 1}" + else: + middle_frames = ['EMPTY'] * num_middle_frames + control_frame_info = "No control frames, all empty" + + # Step 6: Create masks based on mask mode + masks = [1.0] * total_output_frames # 1.0 = inpaint, 0.0 = known + + if when_to_start_masks == "beginning": + # Set known frames (overlap and end) to 0.0, rest stay as 1.0 (inpaint) + for i in range(actual_overlap_frames): + masks[i] = 0.0 # overlap frames are known + for i in range(total_output_frames - num_end_frames, total_output_frames): + masks[i] = 0.0 # end frames are known + mask_info = "Standard: overlap and end frames known, middle frames inpaint" + else: # "after overlap_frames" + # Follow control frame logic + for i in range(actual_overlap_frames): + masks[i] = 0.0 # overlap frames are known + for i in range(total_output_frames - num_end_frames, total_output_frames): + masks[i] = 0.0 # end frames are known + + # For middle section, follow the same logic as control frames + if control_frames and num_middle_frames > 0: + duplicate_count = min(actual_overlap_frames, len(control_frames)) + available_after_dup = len(control_frames) - duplicate_count + if available_after_dup >= num_middle_frames: + # If we have enough control frames after skipping, set those middle frames as known (0.0) + middle_start = actual_overlap_frames + middle_end = middle_start + num_middle_frames + for i in range(middle_start, middle_end): + masks[i] = 0.0 + mask_info = "Follows control logic: overlap, control-covered middle, and end frames known" + else: + mask_info = "Follows control logic: overlap and end frames known, partial middle coverage" + else: + mask_info = "No control frames: only overlap and end frames known" + + # Step 7: Assemble final video + final_video = start_frames + middle_frames + end_frames + + print(f"Control: {control_frame_info}") + print(f"Masks: {mask_info}") + print(f"Final video: {final_video}") + print(f"Masks: {['INPAINT' if m > 0.5 else 'KNOWN' for m in masks]}") + + return { + 'start_frames': start_frames, + 'middle_frames': middle_frames, + 'end_frames': end_frames, + 'final_video': final_video, + 'masks': masks, + 'actual_overlap_frames': actual_overlap_frames, + 'num_middle_frames': num_middle_frames, + 'num_end_frames': num_end_frames, + 'control_frame_info': control_frame_info, + 'mask_info': mask_info + } + +def visualize_comprehensive_simulation(results_list, scenario_names): + """ + Create a comprehensive visual representation showing both frames and masks + """ + fig, axes = plt.subplots(len(results_list), 2, figsize=(20, 4 * len(results_list))) + if len(results_list) == 1: + axes = axes.reshape(1, -1) + + colors = { + 'INPUT': '#87CEEB', # Sky blue + 'CONTROL': '#98FB98', # Pale green + 'END': '#FFA500', # Orange + 'EMPTY': '#D3D3D3' # Light gray + } + + mask_colors = { + 'KNOWN': '#4169E1', # Royal blue + 'INPAINT': '#FF6347' # Tomato red + } + + for i, (results, scenario_name) in enumerate(zip(results_list, scenario_names)): + # Plot frames (left column) + ax_frames = axes[i, 0] + final_video = results['final_video'] + + frame_colors = [] + frame_labels = [] + + for frame in final_video: + if isinstance(frame, str): + if frame == 'EMPTY': + frame_colors.append(colors['EMPTY']) + frame_labels.append('EMPTY') + elif frame == 'END': + frame_colors.append(colors['END']) + frame_labels.append('END') + else: + frame_colors.append(colors['CONTROL']) + frame_labels.append(frame) + else: + # Input frame (number) + frame_colors.append(colors['INPUT']) + frame_labels.append(f'I{frame}') + + x_positions = range(len(final_video)) + bars_frames = ax_frames.bar(x_positions, [1] * len(final_video), + color=frame_colors, edgecolor='black', linewidth=0.5) + + # Add frame labels + for j, (bar, label) in enumerate(zip(bars_frames, frame_labels)): + ax_frames.text(bar.get_x() + bar.get_width()/2, bar.get_height()/2, + label, ha='center', va='center', fontsize=8, rotation=90) + + ax_frames.set_title(f'{scenario_name}\nFrames: {results["control_frame_info"]}') + ax_frames.set_ylabel('Frame Content') + ax_frames.set_xlabel('Frame Position') + ax_frames.set_ylim(0, 1.2) + ax_frames.set_xticks(x_positions) + ax_frames.set_xticklabels([str(i) for i in x_positions]) + + # Plot masks (right column) + ax_masks = axes[i, 1] + masks = results['masks'] + + mask_color_list = [mask_colors['KNOWN'] if m < 0.5 else mask_colors['INPAINT'] for m in masks] + mask_label_list = ['KNOWN' if m < 0.5 else 'INPAINT' for m in masks] + + bars_masks = ax_masks.bar(x_positions, [1] * len(masks), + color=mask_color_list, edgecolor='black', linewidth=0.5) + + # Add mask labels + for j, (bar, label) in enumerate(zip(bars_masks, mask_label_list)): + ax_masks.text(bar.get_x() + bar.get_width()/2, bar.get_height()/2, + label, ha='center', va='center', fontsize=8, rotation=90) + + ax_masks.set_title(f'Masks: {results["mask_info"]}') + ax_masks.set_ylabel('Mask Type') + ax_masks.set_xlabel('Frame Position') + ax_masks.set_ylim(0, 1.2) + ax_masks.set_xticks(x_positions) + ax_masks.set_xticklabels([str(i) for i in x_positions]) + + # Add legends + if i == 0: # Only add legend to first row + # Frame legend + frame_legend_elements = [] + for frame_type, color in colors.items(): + frame_legend_elements.append(plt.Rectangle((0,0),1,1, facecolor=color, label=frame_type)) + ax_frames.legend(handles=frame_legend_elements, loc='upper right', bbox_to_anchor=(1.15, 1)) + + # Mask legend + mask_legend_elements = [] + for mask_type, color in mask_colors.items(): + mask_legend_elements.append(plt.Rectangle((0,0),1,1, facecolor=color, label=mask_type)) + ax_masks.legend(handles=mask_legend_elements, loc='upper right', bbox_to_anchor=(1.15, 1)) + + plt.tight_layout() + plt.savefig('video_continuation_comprehensive_simulation.png', dpi=150, bbox_inches='tight') + plt.show() + +def run_comprehensive_simulation(): + # Test setup + input_frames = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9] # 10 input frames + control_frames = [f'C{i}' for i in range(20)] # 20 control frames (enough for all scenarios) + total_output_frames = 17 # (17-1) % 4 == 0 + overlap_frames = 3 + end_frame = 'END' + + print("=" * 80) + print("VIDEO CONTINUATION GENERATOR COMPREHENSIVE SIMULATION") + print("=" * 80) + print(f"Input frames: {input_frames}") + print(f"Control frames: {control_frames}") + print(f"Total output frames: {total_output_frames}") + print(f"Overlap frames: {overlap_frames}") + print(f"End frame: {end_frame}") + print(f"Middle frames needed: {total_output_frames - overlap_frames - 1} = {total_output_frames - overlap_frames - 1}") + + # Test all combinations + scenarios = [ + ("beginning", "beginning", "Control: Beginning, Masks: Beginning"), + ("beginning", "after overlap_frames", "Control: Beginning, Masks: After Overlap"), + ("after overlap_frames", "beginning", "Control: After Overlap, Masks: Beginning"), + ("after overlap_frames", "after overlap_frames", "Control: After Overlap, Masks: After Overlap"), + ] + + results_list = [] + scenario_names = [] + + for control_mode, mask_mode, display_name in scenarios: + result = simulate_video_continuation_comprehensive( + input_frames, control_frames, total_output_frames, + overlap_frames, control_mode, mask_mode, end_frame + ) + results_list.append(result) + scenario_names.append(display_name) + + # Additional test: Insufficient control frames scenario + print("\n" + "="*50) + print("INSUFFICIENT CONTROL FRAMES TEST") + print("="*50) + + short_control_frames = ['C0', 'C1', 'C2', 'C3', 'C4'] # Only 5 control frames + print(f"Short control frames: {short_control_frames}") + + for control_mode, mask_mode, display_name in scenarios: + result = simulate_video_continuation_comprehensive( + input_frames, short_control_frames, total_output_frames, + overlap_frames, control_mode, mask_mode, end_frame + ) + results_list.append(result) + scenario_names.append(f"{display_name} (Short Control)") + + visualize_comprehensive_simulation(results_list, scenario_names) + +if __name__ == "__main__": + run_comprehensive_simulation() \ No newline at end of file