From 1077d2323c1ac0cca4887c872cae2ddaaf07c629 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 10 Jul 2025 23:41:54 +0300 Subject: [PATCH] fix max res selection --- nodes.py | 87 ++++++++++++++++++++++++++++++++++++++++---------------- 1 file changed, 62 insertions(+), 25 deletions(-) diff --git a/nodes.py b/nodes.py index c9f59d8..ef5cd71 100644 --- a/nodes.py +++ b/nodes.py @@ -622,8 +622,8 @@ class WanVideoImageClipEncode: "clip_vision": ("CLIP_VISION",), "image": ("IMAGE", {"tooltip": "Image to encode"}), "vae": ("WANVAE",), - "generation_width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}), - "generation_height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}), + "generation_width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}), + "generation_height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), }, "optional": { @@ -745,8 +745,8 @@ class WanVideoImageResizeToClosest: def INPUT_TYPES(s): return {"required": { "image": ("IMAGE", {"tooltip": "Image to resize"}), - "generation_width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}), - "generation_height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}), + "generation_width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}), + "generation_height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), "aspect_ratio_preservation": (["keep_input", "stretch_to_new", "crop_to_new"],), }, } @@ -924,8 +924,8 @@ class WanVideoImageToVideoEncode: def INPUT_TYPES(s): return {"required": { "vae": ("WANVAE",), - "width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}), - "height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}), + "width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}), + "height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), "noise_aug_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Strength of noise augmentation, helpful for I2V where some noise can add motion and give sharper results"}), "start_latent_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional latent multiplier, helpful for I2V where lower values allow for more motion"}), @@ -1083,8 +1083,8 @@ class WanVideoEmptyEmbeds: @classmethod def INPUT_TYPES(s): return {"required": { - "width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}), - "height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}), + "width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}), + "height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), }, "optional": { @@ -1114,8 +1114,8 @@ class WanVideoMiniMaxRemoverEmbeds: @classmethod def INPUT_TYPES(s): return {"required": { - "width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}), - "height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}), + "width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}), + "height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), "latents": ("LATENT", {"tooltip": "Encoded latents to use as control signals"}), "mask_latents": ("LATENT", {"tooltip": "Encoded latents to use as mask"}), @@ -1276,8 +1276,8 @@ class WanVideoVACEEncode: def INPUT_TYPES(s): return {"required": { "vae": ("WANVAE",), - "width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}), - "height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}), + "width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}), + "height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}), "vace_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply VACE"}), @@ -1584,6 +1584,9 @@ class WanVideoContextOptions: "freenoise": ("BOOLEAN", {"default": True, "tooltip": "Shuffle the noise"}), "verbose": ("BOOLEAN", {"default": False, "tooltip": "Print debug output"}), }, + "optional": { + "fuse_method": (["linear", "pyramid"], {"default": "linear", "tooltip": "Window weight function: linear=ramps at edges only, pyramid=triangular weights peaking in middle"}), + } } RETURN_TYPES = ("WANVIDCONTEXT", ) @@ -1592,7 +1595,7 @@ class WanVideoContextOptions: CATEGORY = "WanVideoWrapper" DESCRIPTION = "Context options for WanVideo, allows splitting the video into context windows and attemps blending them for longer generations than the model and memory otherwise would allow." - def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise, verbose, image_cond_start_step=6, image_cond_window_count=2, vae=None): + def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise, verbose, image_cond_start_step=6, image_cond_window_count=2, vae=None, fuse_method="linear"): context_options = { "context_schedule":context_schedule, "context_frames":context_frames, @@ -1600,6 +1603,7 @@ class WanVideoContextOptions: "context_overlap":context_overlap, "freenoise":freenoise, "verbose":verbose, + "fuse_method":fuse_method } return (context_options,) @@ -2180,21 +2184,54 @@ class WanVideoSampler: is_looped = False if context_options is not None: - def create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=False): + def create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=False, window_type="linear"): window_mask = torch.ones_like(noise_pred_context) - # Apply left-side blending for all except first chunk (or always in loop mode) - if min(c) > 0 or (looped and max(c) == latent_video_length - 1): - ramp_up = torch.linspace(0, 1, context_overlap, device=noise_pred_context.device) - ramp_up = ramp_up.view(1, -1, 1, 1) - window_mask[:, :context_overlap] = ramp_up + if window_type == "pyramid": + # Create pyramid weights that peak in the middle + length = noise_pred_context.shape[1] + if length % 2 == 0: + max_weight = length // 2 + weight_sequence = list(range(1, max_weight + 1, 1)) + list(range(max_weight, 0, -1)) + else: + max_weight = (length + 1) // 2 + weight_sequence = list(range(1, max_weight, 1)) + [max_weight] + list(range(max_weight - 1, 0, -1)) - # Apply right-side blending for all except last chunk (or always in loop mode) - if max(c) < latent_video_length - 1 or (looped and min(c) == 0): - ramp_down = torch.linspace(1, 0, context_overlap, device=noise_pred_context.device) - ramp_down = ramp_down.view(1, -1, 1, 1) - window_mask[:, -context_overlap:] = ramp_down + # Normalize weights to range from 0 to 1 + max_val = max(weight_sequence) + weight_sequence = [w / max_val for w in weight_sequence] + # Apply the weights to create the mask + weights_tensor = torch.tensor(weight_sequence, device=noise_pred_context.device) + weights_tensor = weights_tensor.view(1, -1, 1, 1) + window_mask = weights_tensor.expand_as(window_mask).clone() + + # Adjust for position in sequence if needed + if not looped: + if min(c) == 0: # First chunk + left_ramp = torch.linspace(0, 1, context_overlap, device=noise_pred_context.device).view(1, -1, 1, 1) + # Clone to avoid in-place memory conflict + left_section = window_mask[:, :context_overlap].clone() + window_mask[:, :context_overlap] = torch.maximum(left_section, left_ramp) + + if max(c) == latent_video_length - 1: # Last chunk + right_ramp = torch.linspace(1, 0, context_overlap, device=noise_pred_context.device).view(1, -1, 1, 1) + # Clone to avoid in-place memory conflict + right_section = window_mask[:, -context_overlap:].clone() + window_mask[:, -context_overlap:] = torch.maximum(right_section, right_ramp) + else: # Original "linear" window masking + # Apply left-side blending for all except first chunk (or always in loop mode) + if min(c) > 0 or (looped and max(c) == latent_video_length - 1): + ramp_up = torch.linspace(0, 1, context_overlap, device=noise_pred_context.device) + ramp_up = ramp_up.view(1, -1, 1, 1) + window_mask[:, :context_overlap] = ramp_up + + # Apply right-side blending for all except last chunk (or always in loop mode) + if max(c) < latent_video_length - 1 or (looped and min(c) == 0): + ramp_down = torch.linspace(1, 0, context_overlap, device=noise_pred_context.device) + ramp_down = ramp_down.view(1, -1, 1, 1) + window_mask[:, -context_overlap:] = ramp_down + return window_mask context_schedule = context_options["context_schedule"] @@ -3057,7 +3094,7 @@ class WanVideoSampler: if cache_args is not None: self.window_tracker.cache_states[window_id] = new_teacache - window_mask = create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=is_looped) + window_mask = create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=is_looped, window_type=context_options["fuse_method"]) noise_pred[:, c] += noise_pred_context * window_mask counter[:, c] += window_mask context_pbar.update_absolute(step_start_progress + (i + 1) * fraction_per_context, steps)