fix max res selection
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user