From 6b9565ed6307168c45770ec08517d7df65b61e3f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 18 Jul 2025 16:15:50 +0300 Subject: [PATCH] Code refactoring --- __init__.py | 12 +- cache_methods/cache_methods.py | 158 ++++ cache_methods/nodes_cache.py | 140 ++++ context.py => context_windows/context.py | 75 +- nodes.py | 875 ++++------------------- nodes_utility.py | 242 +++++++ wanvideo/modules/model.py | 108 +-- wanvideo/schedulers/__init__.py | 85 +++ 8 files changed, 843 insertions(+), 852 deletions(-) create mode 100644 cache_methods/cache_methods.py create mode 100644 cache_methods/nodes_cache.py rename context.py => context_windows/context.py (63%) create mode 100644 nodes_utility.py diff --git a/__init__.py b/__init__.py index 04b0240..cd84e62 100644 --- a/__init__.py +++ b/__init__.py @@ -8,7 +8,8 @@ from .controlnet.nodes import NODE_CLASS_MAPPINGS as CONTROLNET_NODE_CLASS_MAPPI from .ATI.nodes import NODE_CLASS_MAPPINGS as ATI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as ATI_NODE_DISPLAY_NAME_MAPPINGS from .multitalk.nodes import NODE_CLASS_MAPPINGS as MULTITALK_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MULTITALK_NODE_DISPLAY_NAME_MAPPINGS from .nodes_model_loading import NODE_CLASS_MAPPINGS as MODEL_LOADING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS - +from .nodes_utility import NODE_CLASS_MAPPINGS as UTILITY_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UTILITY_NODE_DISPLAY_NAME_MAPPINGS +from .cache_methods.nodes_cache import NODE_CLASS_MAPPINGS as NODE_CACHE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as NODE_CACHE_DISPLAY_NAME_MAPPINGS try: from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS except ImportError: @@ -16,8 +17,6 @@ except ImportError: UNIANIMATE_NODE_CLASS_MAPPINGS = {} UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS = {} -#from .causvid.nodes import NODE_CLASS_MAPPINGS as CAUSVID_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CAUSVID_NODE_DISPLAY_NAME_MAPPINGS - NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS) @@ -28,10 +27,10 @@ NODE_CLASS_MAPPINGS.update(CONTROLNET_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(ATI_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(MULTITALK_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(MODEL_LOADING_NODE_CLASS_MAPPINGS) +NODE_CLASS_MAPPINGS.update(UTILITY_NODE_CLASS_MAPPINGS) +NODE_CLASS_MAPPINGS.update(NODE_CACHE_CLASS_MAPPINGS) -#NODE_CLASS_MAPPINGS.update(CAUSVID_NODE_CLASS_MAPPINGS) - NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(SKYREELS_NODE_DISPLAY_NAME_MAPPINGS) @@ -42,7 +41,8 @@ NODE_DISPLAY_NAME_MAPPINGS.update(CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(ATI_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(MULTITALK_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS) +NODE_DISPLAY_NAME_MAPPINGS.update(UTILITY_NODE_DISPLAY_NAME_MAPPINGS) +NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_DISPLAY_NAME_MAPPINGS) -#NODE_DISPLAY_NAME_MAPPINGS.update(CAUSVID_NODE_DISPLAY_NAME_MAPPINGS) __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/cache_methods/cache_methods.py b/cache_methods/cache_methods.py new file mode 100644 index 0000000..1b7bb12 --- /dev/null +++ b/cache_methods/cache_methods.py @@ -0,0 +1,158 @@ +from ..utils import log +import torch + +def set_transformer_cache_method(transformer, timesteps, cache_args=None): + transformer.cache_device = cache_args["cache_device"] + if cache_args["cache_type"] == "TeaCache": + log.info(f"TeaCache: Using cache device: {transformer.cache_device}") + transformer.teacache_state.clear_all() + transformer.enable_teacache = True + transformer.rel_l1_thresh = cache_args["rel_l1_thresh"] + transformer.teacache_start_step = cache_args["start_step"] + transformer.teacache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] + transformer.teacache_use_coefficients = cache_args["use_coefficients"] + transformer.teacache_mode = cache_args["mode"] + elif cache_args["cache_type"] == "MagCache": + log.info(f"MagCache: Using cache device: {transformer.cache_device}") + transformer.magcache_state.clear_all() + transformer.enable_magcache = True + transformer.magcache_start_step = cache_args["start_step"] + transformer.magcache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] + transformer.magcache_thresh = cache_args["magcache_thresh"] + transformer.magcache_K = cache_args["magcache_K"] + elif cache_args["cache_type"] == "EasyCache": + log.info(f"EasyCache: Using cache device: {transformer.cache_device}") + transformer.easycache_state.clear_all() + transformer.enable_easycache = True + transformer.easycache_start_step = cache_args["start_step"] + transformer.easycache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] + transformer.easycache_thresh = cache_args["easycache_thresh"] + return transformer + +class TeaCacheState: + def __init__(self, cache_device='cpu'): + self.cache_device = cache_device + self.states = {} + self._next_pred_id = 0 + + def new_prediction(self, cache_device='cpu'): + """Create new prediction state and return its ID""" + self.cache_device = cache_device + pred_id = self._next_pred_id + self._next_pred_id += 1 + self.states[pred_id] = { + 'previous_residual': None, + 'accumulated_rel_l1_distance': 0, + 'previous_modulated_input': None, + 'skipped_steps': [], + } + return pred_id + + def update(self, pred_id, **kwargs): + """Update state for specific prediction""" + if pred_id not in self.states: + return None + for key, value in kwargs.items(): + self.states[pred_id][key] = value + + def get(self, pred_id): + return self.states.get(pred_id, {}) + + def clear_all(self): + self.states = {} + self._next_pred_id = 0 + +class MagCacheState: + def __init__(self, cache_device='cpu'): + self.cache_device = cache_device + self.states = {} + self._next_pred_id = 0 + + def new_prediction(self, cache_device='cpu'): + """Create new prediction state and return its ID""" + self.cache_device = cache_device + pred_id = self._next_pred_id + self._next_pred_id += 1 + self.states[pred_id] = { + 'residual_cache': None, + 'accumulated_ratio': 1.0, + 'accumulated_steps': 0, + 'accumulated_err': 0, + 'skipped_steps': [], + } + return pred_id + + def update(self, pred_id, **kwargs): + """Update state for specific prediction""" + if pred_id not in self.states: + return None + for key, value in kwargs.items(): + self.states[pred_id][key] = value + + def get(self, pred_id): + return self.states.get(pred_id, {}) + + def clear_all(self): + self.states = {} + self._next_pred_id = 0 + +class EasyCacheState: + def __init__(self, cache_device='cpu'): + self.cache_device = cache_device + self.states = {} + self._next_pred_id = 0 + + def new_prediction(self, cache_device='cpu'): + """Create a new prediction state and return its ID.""" + self.cache_device = cache_device + pred_id = self._next_pred_id + self._next_pred_id += 1 + self.states[pred_id] = { + 'previous_raw_input': None, + 'previous_raw_output': None, + 'cache': None, + 'accumulated_error': 0.0, + 'skipped_steps': [], + } + return pred_id + + def update(self, pred_id, **kwargs): + """Update state for a specific prediction.""" + if pred_id not in self.states: + return None + for key, value in kwargs.items(): + self.states[pred_id][key] = value + + def get(self, pred_id): + return self.states.get(pred_id, {}) + + def clear_all(self): + self.states = {} + self._next_pred_id = 0 + +def relative_l1_distance(last_tensor, current_tensor): + l1_distance = torch.abs(last_tensor.to(current_tensor.device) - current_tensor).mean() + norm = torch.abs(last_tensor).mean() + relative_l1_distance = l1_distance / norm + return relative_l1_distance.to(torch.float32).to(current_tensor.device) + +def cache_report(transformer, cache_args): + cache_type = cache_args["cache_type"] + states = ( + transformer.teacache_state.states if cache_type == "TeaCache" else + transformer.magcache_state.states if cache_type == "MagCache" else + transformer.easycache_state.states if cache_type == "EasyCache" else + None + ) + state_names = { + 0: "conditional", + 1: "unconditional" + } + for pred_id, state in states.items(): + name = state_names.get(pred_id, f"prediction_{pred_id}") + if 'skipped_steps' in state: + log.info(f"{cache_type} skipped: {len(state['skipped_steps'])} {name} steps: {state['skipped_steps']}") + transformer.teacache_state.clear_all() + transformer.magcache_state.clear_all() + transformer.easycache_state.clear_all() + del states \ No newline at end of file diff --git a/cache_methods/nodes_cache.py b/cache_methods/nodes_cache.py new file mode 100644 index 0000000..f9402b6 --- /dev/null +++ b/cache_methods/nodes_cache.py @@ -0,0 +1,140 @@ +from comfy import model_management as mm + +class WanVideoTeaCache: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "rel_l1_thresh": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.001, + "tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts. Good value range for 1.3B: 0.05 - 0.08, for other models 0.15-0.30"}), + "start_step": ("INT", {"default": 1, "min": 0, "max": 9999, "step": 1, "tooltip": "Start percentage of the steps to apply TeaCache"}), + "end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "End steps to apply TeaCache"}), + "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), + "use_coefficients": ("BOOLEAN", {"default": True, "tooltip": "Use calculated coefficients for more accuracy. When enabled therel_l1_thresh should be about 10 times higher than without"}), + }, + "optional": { + "mode": (["e", "e0"], {"default": "e", "tooltip": "Choice between using e (time embeds, default) or e0 (modulated time embeds)"}), + }, + } + RETURN_TYPES = ("CACHEARGS",) + RETURN_NAMES = ("cache_args",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = """ +Patch WanVideo model to use TeaCache. Speeds up inference by caching the output and +applying it instead of doing the step. Best results are achieved by choosing the +appropriate coefficients for the model. Early steps should never be skipped, with too +aggressive values this can happen and the motion suffers. Starting later can help with that too. +When NOT using coefficients, the threshold value should be +about 10 times smaller than the value used with coefficients. + +Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4Wan2.1: + + +
++-------------------+--------+---------+--------+
+|       Model       |  Low   | Medium  |  High  |
++-------------------+--------+---------+--------+
+| Wan2.1 t2v 1.3B  |  0.05  |  0.07   |  0.08  |
+| Wan2.1 t2v 14B   |  0.14  |  0.15   |  0.20  |
+| Wan2.1 i2v 480P  |  0.13  |  0.19   |  0.26  |
+| Wan2.1 i2v 720P  |  0.18  |  0.20   |  0.30  |
++-------------------+--------+---------+--------+
+
+""" + + def process(self, rel_l1_thresh, start_step, end_step, cache_device, use_coefficients, mode="e"): + if cache_device == "main_device": + cache_device = mm.get_torch_device() + else: + cache_device = mm.unet_offload_device() + cache_args = { + "cache_type": "TeaCache", + "rel_l1_thresh": rel_l1_thresh, + "start_step": start_step, + "end_step": end_step, + "cache_device": cache_device, + "use_coefficients": use_coefficients, + "mode": mode, + } + return (cache_args,) + +class WanVideoMagCache: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "magcache_thresh": ("FLOAT", {"default": 0.02, "min": 0.0, "max": 0.3, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}), + "magcache_K": ("INT", {"default": 4, "min": 0, "max": 6, "step": 1, "tooltip": "The maxium skip steps of MagCache."}), + "start_step": ("INT", {"default": 1, "min": 0, "max": 9999, "step": 1, "tooltip": "Step to start applying MagCache"}), + "end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying MagCache"}), + "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), + }, + } + RETURN_TYPES = ("CACHEARGS",) + RETURN_NAMES = ("cache_args",) + FUNCTION = "setargs" + CATEGORY = "WanVideoWrapper" + EXPERIMENTAL = True + DESCRIPTION = "MagCache for WanVideoWrapper, source https://github.com/Zehong-Ma/MagCache" + + def setargs(self, magcache_thresh, magcache_K, start_step, end_step, cache_device): + if cache_device == "main_device": + cache_device = mm.get_torch_device() + else: + cache_device = mm.unet_offload_device() + + cache_args = { + "cache_type": "MagCache", + "magcache_thresh": magcache_thresh, + "magcache_K": magcache_K, + "start_step": start_step, + "end_step": end_step, + "cache_device": cache_device, + } + return (cache_args,) + +class WanVideoEasyCache: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "easycache_thresh": ("FLOAT", {"default": 0.015, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}), + "start_step": ("INT", {"default": 10, "min": 1, "max": 9999, "step": 1, "tooltip": "Step to start applying EasyCache"}), + "end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying EasyCache"}), + "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), + }, + } + RETURN_TYPES = ("CACHEARGS",) + RETURN_NAMES = ("cache_args",) + FUNCTION = "setargs" + CATEGORY = "WanVideoWrapper" + EXPERIMENTAL = True + DESCRIPTION = "EasyCache for WanVideoWrapper, source https://github.com/H-EmbodVis/EasyCache" + + def setargs(self, easycache_thresh, start_step, end_step, cache_device): + if cache_device == "main_device": + cache_device = mm.get_torch_device() + else: + cache_device = mm.unet_offload_device() + + cache_args = { + "cache_type": "EasyCache", + "easycache_thresh": easycache_thresh, + "start_step": start_step, + "end_step": end_step, + "cache_device": cache_device, + } + return (cache_args,) + + +NODE_CLASS_MAPPINGS = { + "WanVideoTeaCache": WanVideoTeaCache, + "WanVideoMagCache": WanVideoMagCache, + "WanVideoEasyCache": WanVideoEasyCache, + } +NODE_DISPLAY_NAME_MAPPINGS = { + "WanVideoTeaCache": "WaWanVideo TeaCache", + "WanVideoMagCache": "WanVideo MagCache", + "WanVideoEasyCache": "WanVideo EasyCache" + } \ No newline at end of file diff --git a/context.py b/context_windows/context.py similarity index 63% rename from context.py rename to context_windows/context.py index 6a30fed..867458a 100644 --- a/context.py +++ b/context_windows/context.py @@ -1,6 +1,6 @@ import numpy as np from typing import Callable, Optional, List - +import torch def ordered_halving(val): bin_str = f"{val:064b}" @@ -182,3 +182,76 @@ def get_total_steps( ) for i in range(len(timesteps)) ) + +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) + + 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)) + + # 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 + +class WindowTracker: + def __init__(self, verbose=False): + self.window_map = {} # Maps frame sequence to persistent ID + self.next_id = 0 + self.cache_states = {} # Maps persistent ID to teacache state + self.verbose = verbose + + def get_window_id(self, frames): + key = tuple(sorted(frames)) # Order-independent frame sequence + if key not in self.window_map: + self.window_map[key] = self.next_id + if self.verbose: + log.info(f"New window pattern {key} -> ID {self.next_id}") + self.next_id += 1 + return self.window_map[key] + + def get_teacache(self, window_id, base_state): + if window_id not in self.cache_states: + if self.verbose: + log.info(f"Initializing persistent teacache for window {window_id}") + self.cache_states[window_id] = base_state.copy() + return self.cache_states[window_id] diff --git a/nodes.py b/nodes.py index ef23bd1..3df8cb9 100644 --- a/nodes.py +++ b/nodes.py @@ -1,39 +1,36 @@ -import os +import os, gc, math import torch import torch.nn.functional as F -import gc -from .utils import log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, is_image_black import numpy as np -import math from tqdm import tqdm +import inspect + +from diffusers.schedulers import FlowMatchEulerDiscreteScheduler from .wanvideo.modules.clip import CLIPModel from .wanvideo.modules.model import rope_params from .wanvideo.modules.t5 import T5EncoderModel - -from .wanvideo.schedulers import ( - FlowDPMSolverMultistepScheduler, FlowUniPCMultistepScheduler, - FlowMatchScheduler, FlowMatchSchedulerPusa, FlowMatchLCMScheduler, - get_sampling_sigmas, retrieve_timesteps -) - -from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler +from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps from .multitalk.multitalk import timestep_transform, add_noise - +from .utils import log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, is_image_black +from .cache_methods.cache_methods import cache_report from .enhance_a_video.globals import set_enhance_weight, set_num_frames from .taehv import TAEHV from einops import rearrange import folder_paths -import comfy.model_management as mm +from comfy import model_management as mm from comfy.utils import load_torch_file, ProgressBar, common_upscale from comfy.clip_vision import clip_preprocess, ClipVisionModel from comfy.cli_args import args, LatentPreviewMethod script_directory = os.path.dirname(os.path.abspath(__file__)) +device = mm.get_torch_device() +offload_device = mm.unet_offload_device() + VAE_STRIDE = (4, 8, 8) PATCH_SIZE = (1, 2, 2) @@ -97,132 +94,6 @@ class WanVideoVRAMManagement: def setargs(self, **kwargs): return (kwargs, ) -class WanVideoTeaCache: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "rel_l1_thresh": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.001, - "tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts. Good value range for 1.3B: 0.05 - 0.08, for other models 0.15-0.30"}), - "start_step": ("INT", {"default": 1, "min": 0, "max": 9999, "step": 1, "tooltip": "Start percentage of the steps to apply TeaCache"}), - "end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "End steps to apply TeaCache"}), - "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), - "use_coefficients": ("BOOLEAN", {"default": True, "tooltip": "Use calculated coefficients for more accuracy. When enabled therel_l1_thresh should be about 10 times higher than without"}), - }, - "optional": { - "mode": (["e", "e0"], {"default": "e", "tooltip": "Choice between using e (time embeds, default) or e0 (modulated time embeds)"}), - }, - } - RETURN_TYPES = ("CACHEARGS",) - RETURN_NAMES = ("cache_args",) - FUNCTION = "process" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = """ -Patch WanVideo model to use TeaCache. Speeds up inference by caching the output and -applying it instead of doing the step. Best results are achieved by choosing the -appropriate coefficients for the model. Early steps should never be skipped, with too -aggressive values this can happen and the motion suffers. Starting later can help with that too. -When NOT using coefficients, the threshold value should be -about 10 times smaller than the value used with coefficients. - -Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4Wan2.1: - - -
-+-------------------+--------+---------+--------+
-|       Model       |  Low   | Medium  |  High  |
-+-------------------+--------+---------+--------+
-| Wan2.1 t2v 1.3B  |  0.05  |  0.07   |  0.08  |
-| Wan2.1 t2v 14B   |  0.14  |  0.15   |  0.20  |
-| Wan2.1 i2v 480P  |  0.13  |  0.19   |  0.26  |
-| Wan2.1 i2v 720P  |  0.18  |  0.20   |  0.30  |
-+-------------------+--------+---------+--------+
-
-""" - - def process(self, rel_l1_thresh, start_step, end_step, cache_device, use_coefficients, mode="e"): - if cache_device == "main_device": - cache_device = mm.get_torch_device() - else: - cache_device = mm.unet_offload_device() - cache_args = { - "cache_type": "TeaCache", - "rel_l1_thresh": rel_l1_thresh, - "start_step": start_step, - "end_step": end_step, - "cache_device": cache_device, - "use_coefficients": use_coefficients, - "mode": mode, - } - return (cache_args,) - -class WanVideoMagCache: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "magcache_thresh": ("FLOAT", {"default": 0.02, "min": 0.0, "max": 0.3, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}), - "magcache_K": ("INT", {"default": 4, "min": 0, "max": 6, "step": 1, "tooltip": "The maxium skip steps of MagCache."}), - "start_step": ("INT", {"default": 1, "min": 0, "max": 9999, "step": 1, "tooltip": "Step to start applying MagCache"}), - "end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying MagCache"}), - "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), - }, - } - RETURN_TYPES = ("CACHEARGS",) - RETURN_NAMES = ("cache_args",) - FUNCTION = "setargs" - CATEGORY = "WanVideoWrapper" - EXPERIMENTAL = True - DESCRIPTION = "MagCache for WanVideoWrapper, source https://github.com/Zehong-Ma/MagCache" - - def setargs(self, magcache_thresh, magcache_K, start_step, end_step, cache_device): - if cache_device == "main_device": - cache_device = mm.get_torch_device() - else: - cache_device = mm.unet_offload_device() - - cache_args = { - "cache_type": "MagCache", - "magcache_thresh": magcache_thresh, - "magcache_K": magcache_K, - "start_step": start_step, - "end_step": end_step, - "cache_device": cache_device, - } - return (cache_args,) - -class WanVideoEasyCache: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "easycache_thresh": ("FLOAT", {"default": 0.015, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}), - "start_step": ("INT", {"default": 10, "min": 1, "max": 9999, "step": 1, "tooltip": "Step to start applying EasyCache"}), - "end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying EasyCache"}), - "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), - }, - } - RETURN_TYPES = ("CACHEARGS",) - RETURN_NAMES = ("cache_args",) - FUNCTION = "setargs" - CATEGORY = "WanVideoWrapper" - EXPERIMENTAL = True - DESCRIPTION = "EasyCache for WanVideoWrapper, source https://github.com/H-EmbodVis/EasyCache" - - def setargs(self, easycache_thresh, start_step, end_step, cache_device): - if cache_device == "main_device": - cache_device = mm.get_torch_device() - else: - cache_device = mm.unet_offload_device() - - cache_args = { - "cache_type": "EasyCache", - "easycache_thresh": easycache_thresh, - "start_step": start_step, - "end_step": end_step, - "cache_device": cache_device, - } - return (cache_args,) class WanVideoEnhanceAVideo: @classmethod @@ -370,10 +241,6 @@ class LoadWanVideoT5TextEncoder: DESCRIPTION = "Loads Wan text_encoder model from 'ComfyUI/models/LLM'" def loadmodel(self, model_name, precision, load_device="offload_device", quantization="disabled"): - - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() - text_encoder_load_device = device if load_device == "main_device" else offload_device tokenizer_path = os.path.join(script_directory, "configs", "T5_tokenizer") @@ -478,10 +345,6 @@ class LoadWanVideoClipTextEncoder: DESCRIPTION = "Loads Wan clip_vision model from 'ComfyUI/models/clip_vision'" def loadmodel(self, model_name, precision, load_device="offload_device"): - - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() - text_encoder_load_device = device if load_device == "main_device" else offload_device dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] @@ -522,10 +385,6 @@ class WanVideoTextEncode: DESCRIPTION = "Encodes text prompts into text embeddings. For rudimentary prompt travel you can input multiple prompts separated by '|', they will be equally spread over the video length" def process(self, t5, positive_prompt, negative_prompt,force_offload=True, model_to_offload=None): - - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() - if model_to_offload is not None: log.info(f"Moving video model to {offload_device}") model_to_offload.model.to(offload_device) @@ -607,10 +466,6 @@ class WanVideoTextEncodeSingle: DESCRIPTION = "Encodes text prompt into text embedding." def process(self, t5, prompt, force_offload=True, model_to_offload=None): - - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() - if model_to_offload is not None: log.info(f"Moving video model to {offload_device}") model_to_offload.model.to(offload_device) @@ -682,7 +537,6 @@ class WanVideoTextEmbedBridge: DESCRIPTION = "Bridge between ComfyUI native text embedding and WanVideoWrapper text embedding" def process(self, positive, negative=None): - device=mm.get_torch_device() prompt_embeds_dict = { "prompt_embeds": positive[0][0].to(device), "negative_prompt_embeds": negative[0][0].to(device) if negative is not None else None, @@ -720,9 +574,6 @@ class WanVideoImageClipEncode: def process(self, clip_vision, vae, image, num_frames, generation_width, generation_height, force_offload=True, noise_aug_strength=0.0, latent_strength=1.0, clip_embed_strength=1.0, adjust_resolution=True): - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() - self.image_mean = [0.48145466, 0.4578275, 0.40821073] self.image_std = [0.26862954, 0.26130258, 0.27577711] @@ -814,50 +665,6 @@ class WanVideoImageClipEncode: } return (image_embeds,) - -class WanVideoImageResizeToClosest: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "image": ("IMAGE", {"tooltip": "Image to resize"}), - "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"],), - }, - } - - RETURN_TYPES = ("IMAGE", "INT", "INT", ) - RETURN_NAMES = ("image","width","height",) - FUNCTION = "process" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Resizes image to the closest supported resolution based on aspect ratio and max pixels, according to the original code" - - def process(self, image, generation_width, generation_height, aspect_ratio_preservation ): - - H, W = image.shape[1], image.shape[2] - max_area = generation_width * generation_height - - crop = "disabled" - - if aspect_ratio_preservation == "keep_input": - aspect_ratio = H / W - elif aspect_ratio_preservation == "stretch_to_new" or aspect_ratio_preservation == "crop_to_new": - aspect_ratio = generation_height / generation_width - if aspect_ratio_preservation == "crop_to_new": - crop = "center" - - lat_h = round( - np.sqrt(max_area * aspect_ratio) // VAE_STRIDE[1] // - PATCH_SIZE[1] * PATCH_SIZE[1]) - lat_w = round( - np.sqrt(max_area / aspect_ratio) // VAE_STRIDE[2] // - PATCH_SIZE[2] * PATCH_SIZE[2]) - h = lat_h * VAE_STRIDE[1] - w = lat_w * VAE_STRIDE[2] - - resized_image = common_upscale(image.movedim(-1, 1), w, h, "lanczos", crop).movedim(1, -1) - - return (resized_image, w, h) #region clip vision class WanVideoClipVisionEncode: @@ -886,10 +693,6 @@ class WanVideoClipVisionEncode: CATEGORY = "WanVideoWrapper" def process(self, clip_vision, image_1, strength_1, strength_2, force_offload, crop, combine_embeds, image_2=None, negative_image=None, tiles=0, ratio=1.0): - - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() - image_mean = [0.48145466, 0.4578275, 0.40821073] image_std = [0.26862954, 0.26130258, 0.27577711] @@ -1035,9 +838,6 @@ class WanVideoImageToVideoEncode: start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False, temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None): - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() - H = height W = width @@ -1381,10 +1181,7 @@ class WanVideoVACEEncode: CATEGORY = "WanVideoWrapper" def process(self, vae, width, height, num_frames, strength, vace_start_percent, vace_end_percent, input_frames=None, ref_images=None, input_masks=None, prev_vace_embeds=None, tiled_vae=False): - - self.device = mm.get_torch_device() - offload_device = mm.unet_offload_device() - self.vae = vae.to(self.device) + vae = vae.to(device) width = (width // 16) * 16 height = (height // 16) * 16 @@ -1394,19 +1191,19 @@ class WanVideoVACEEncode: width // VAE_STRIDE[2]) # vace context encode if input_frames is None: - input_frames = torch.zeros((1, 3, num_frames, height, width), device=self.device, dtype=self.vae.dtype) + input_frames = torch.zeros((1, 3, num_frames, height, width), device=device, dtype=vae.dtype) else: input_frames = input_frames[:num_frames] input_frames = common_upscale(input_frames.clone().movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1) - input_frames = input_frames.to(self.vae.dtype).to(self.device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W + input_frames = input_frames.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W input_frames = input_frames * 2 - 1 if input_masks is None: - input_masks = torch.ones_like(input_frames, device=self.device) + input_masks = torch.ones_like(input_frames, device=device) else: print("input_masks shape", input_masks.shape) input_masks = input_masks[:num_frames] input_masks = common_upscale(input_masks.clone().unsqueeze(1), width, height, "nearest-exact", "disabled").squeeze(1) - input_masks = input_masks.to(self.vae.dtype).to(self.device) + input_masks = input_masks.to(vae.dtype).to(device) input_masks = input_masks.unsqueeze(-1).unsqueeze(0).permute(0, 4, 1, 2, 3).repeat(1, 3, 1, 1, 1) # B, C, T, H, W if ref_images is not None: @@ -1433,15 +1230,15 @@ class WanVideoVACEEncode: ref_images = padded ref_images = common_upscale(ref_images.movedim(-1, 1), width, height, "lanczos", "center").movedim(1, -1) - ref_images = ref_images.to(self.vae.dtype).to(self.device).unsqueeze(0).permute(0, 4, 1, 2, 3).unsqueeze(0) + ref_images = ref_images.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3).unsqueeze(0) ref_images = ref_images * 2 - 1 - z0 = self.vace_encode_frames(input_frames, ref_images, masks=input_masks, tiled_vae=tiled_vae) - self.vae.model.clear_cache() + z0 = self.vace_encode_frames(vae, input_frames, ref_images, masks=input_masks, tiled_vae=tiled_vae) + vae.model.clear_cache() m0 = self.vace_encode_masks(input_masks, ref_images) z = self.vace_latent(z0, m0) - self.vae.to(offload_device) + vae.to(offload_device) vace_input = { "vace_context": z, @@ -1461,29 +1258,29 @@ class WanVideoVACEEncode: vace_input["additional_vace_inputs"].append(prev_vace_embeds) return (vace_input,) - def vace_encode_frames(self, frames, ref_images, masks=None, tiled_vae=False): + def vace_encode_frames(self, vae, frames, ref_images, masks=None, tiled_vae=False): if ref_images is None: ref_images = [None] * len(frames) else: assert len(frames) == len(ref_images) if masks is None: - latents = self.vae.encode(frames, device=self.device, tiled=tiled_vae) + latents = vae.encode(frames, device=device, tiled=tiled_vae) else: inactive = [i * (1 - m) + 0 * m for i, m in zip(frames, masks)] reactive = [i * m + 0 * (1 - m) for i, m in zip(frames, masks)] - inactive = self.vae.encode(inactive, device=self.device, tiled=tiled_vae) - reactive = self.vae.encode(reactive, device=self.device, tiled=tiled_vae) + inactive = vae.encode(inactive, device=device, tiled=tiled_vae) + reactive = vae.encode(reactive, device=device, tiled=tiled_vae) latents = [torch.cat((u, c), dim=0) for u, c in zip(inactive, reactive)] - self.vae.model.clear_cache() + vae.model.clear_cache() cat_latents = [] for latent, refs in zip(latents, ref_images): if refs is not None: if masks is None: - ref_latent = self.vae.encode(refs, device=self.device, tiled=tiled_vae) + ref_latent = vae.encode(refs, device=device, tiled=tiled_vae) else: print("refs shape", refs.shape)#torch.Size([3, 1, 512, 512]) - ref_latent = self.vae.encode(refs, device=self.device, tiled=tiled_vae) + ref_latent = vae.encode(refs, device=device, tiled=tiled_vae) ref_latent = [torch.cat((u, torch.zeros_like(u)), dim=0) for u in ref_latent] assert all([x.shape[1] == 1 for x in ref_latent]) latent = torch.cat([*ref_latent, latent], dim=1) @@ -1526,134 +1323,6 @@ class WanVideoVACEEncode: def vace_latent(self, z, m): return [torch.cat([zz, mm], dim=0) for zz, mm in zip(z, m)] -class ExtractStartFramesForContinuations: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "input_video_frames": ("IMAGE", {"tooltip": "Input video frames to extract the start frames from."}), - "num_frames": ("INT", {"default": 10, "min": 1, "max": 1024, "step": 1, "tooltip": "Number of frames to get from the start of the video."}), - }, - } - - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("start_frames",) - FUNCTION = "get_start_frames" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Extracts the first N frames from a video sequence for continuations." - - def get_start_frames(self, input_video_frames, num_frames): - if input_video_frames is None or input_video_frames.shape[0] == 0: - log.warning("Input video frames are empty. Returning an empty tensor.") - if input_video_frames is not None: - return (torch.empty((0,) + input_video_frames.shape[1:], dtype=input_video_frames.dtype),) - else: - # Return a tensor with 4 dimensions, as expected for an IMAGE type. - return (torch.empty((0, 64, 64, 3), dtype=torch.float32),) - - total_frames = input_video_frames.shape[0] - num_to_get = min(num_frames, total_frames) - - if num_to_get < num_frames: - log.warning(f"Requested {num_frames} frames, but input video only has {total_frames} frames. Returning first {num_to_get} frames.") - - start_frames = input_video_frames[:num_to_get] - - return (start_frames.cpu().float(),) - -class WanVideoVACEStartToEndFrame: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), - "empty_frame_level": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "White level of empty frame to use"}), - }, - "optional": { - "start_image": ("IMAGE",), - "end_image": ("IMAGE",), - "control_images": ("IMAGE",), - "inpaint_mask": ("MASK", {"tooltip": "Inpaint mask to use for the empty frames"}), - "start_index": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Index to start from"}), - "end_index": ("INT", {"default": -1, "min": -10000, "max": 10000, "step": 1, "tooltip": "Index to end at"}), - }, - } - - RETURN_TYPES = ("IMAGE", "MASK", ) - RETURN_NAMES = ("images", "masks",) - FUNCTION = "process" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Helper node to create start/end frame batch and masks for VACE" - - def process(self, num_frames, empty_frame_level, start_image=None, end_image=None, control_images=None, inpaint_mask=None, start_index=0, end_index=-1): - - B, H, W, C = start_image.shape if start_image is not None else end_image.shape - device = start_image.device if start_image is not None else end_image.device - - # Convert negative end_index to positive - if end_index < 0: - end_index = num_frames + end_index - - # Create output batch with empty frames - out_batch = torch.ones((num_frames, H, W, 3), device=device) * empty_frame_level - - # Create mask tensor with proper dimensions - masks = torch.ones((num_frames, H, W), device=device) - - # Pre-process all images at once to avoid redundant work - if end_image is not None and (end_image.shape[1] != H or end_image.shape[2] != W): - end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1) - - if control_images is not None and (control_images.shape[1] != H or control_images.shape[2] != W): - control_images = common_upscale(control_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1) - - # Place start image at start_index - if start_image is not None: - frames_to_copy = min(start_image.shape[0], num_frames - start_index) - if frames_to_copy > 0: - out_batch[start_index:start_index + frames_to_copy] = start_image[:frames_to_copy] - masks[start_index:start_index + frames_to_copy] = 0 - - # Place end image at end_index - if end_image is not None: - # Calculate where to start placing end images - end_start = end_index - end_image.shape[0] + 1 - if end_start < 0: # Handle case where end images won't all fit - end_image = end_image[abs(end_start):] - end_start = 0 - - frames_to_copy = min(end_image.shape[0], num_frames - end_start) - if frames_to_copy > 0: - out_batch[end_start:end_start + frames_to_copy] = end_image[:frames_to_copy] - masks[end_start:end_start + frames_to_copy] = 0 - - # Apply control images to remaining frames that don't have start or end images - if control_images is not None: - # Create a mask of frames that are still empty (mask == 1) - empty_frames = masks.sum(dim=(1, 2)) > 0.5 * H * W - - if empty_frames.any(): - # Only apply control images where they exist - control_length = control_images.shape[0] - for frame_idx in range(num_frames): - if empty_frames[frame_idx] and frame_idx < control_length: - out_batch[frame_idx] = control_images[frame_idx] - - # Apply inpaint mask if provided - if inpaint_mask is not None: - inpaint_mask = common_upscale(inpaint_mask.unsqueeze(1), W, H, "nearest-exact", "disabled").squeeze(1).to(device) - - # Handle different mask lengths efficiently - if inpaint_mask.shape[0] > num_frames: - inpaint_mask = inpaint_mask[:num_frames] - elif inpaint_mask.shape[0] < num_frames: - repeat_factor = (num_frames + inpaint_mask.shape[0] - 1) // inpaint_mask.shape[0] # Ceiling division - inpaint_mask = inpaint_mask.repeat(repeat_factor, 1, 1)[:num_frames] - - # Apply mask in one operation - masks = inpaint_mask * masks - - return (out_batch.cpu().float(), masks.cpu().float()) - #region context options class WanVideoContextOptions: @@ -1691,55 +1360,6 @@ class WanVideoContextOptions: return (context_options,) -class CreateCFGScheduleFloatList: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "steps": ("INT", {"default": 30, "min": 2, "max": 1000, "step": 1, "tooltip": "Number of steps to schedule cfg for"} ), - "cfg_scale_start": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 30.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}), - "cfg_scale_end": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 30.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}), - "interpolation": (["linear", "ease_in", "ease_out"], {"default": "linear", "tooltip": "Interpolation method to use for the cfg scale"}), - "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "Start percent of the steps to apply cfg"}), - "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "End percent of the steps to apply cfg"}), - } - } - - RETURN_TYPES = ("FLOAT", ) - RETURN_NAMES = ("float_list",) - FUNCTION = "process" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Helper node to generate a list of floats that can be used to schedule cfg scale for the steps, outside the set range cfg is set to 1.0" - - def process(self, steps, cfg_scale_start, cfg_scale_end, interpolation, start_percent, end_percent): - - # Create a list of floats for the cfg schedule - cfg_list = [1.0] * steps - start_idx = min(int(steps * start_percent), steps - 1) - end_idx = min(int(steps * end_percent), steps - 1) - - for i in range(start_idx, end_idx + 1): - if i >= steps: - break - - if end_idx == start_idx: - t = 0 - else: - t = (i - start_idx) / (end_idx - start_idx) - - if interpolation == "linear": - factor = t - elif interpolation == "ease_in": - factor = t * t - elif interpolation == "ease_out": - factor = t * (2 - t) - - cfg_list[i] = round(cfg_scale_start + factor * (cfg_scale_end - cfg_scale_start), 2) - - # If start_percent > 0, always include the first step - if start_percent > 0: - cfg_list[0] = 1.0 - - return (cfg_list,) class WanVideoFlowEdit: @classmethod @@ -1851,8 +1471,6 @@ class WanVideoSampler: "default": 'unipc' }), "riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 6. Allows for new frames to be generated after without looping"}), - - }, "optional": { "text_embeds": ("WANVIDEOTEXTEMBEDS", ), @@ -1897,9 +1515,6 @@ class WanVideoSampler: raise Exception("multitalk scheduler is only for multitalk sampling when using ImagetoVideoMultiTalk -node") transformer_options = patcher.model_options.get("transformer_options", None) - - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() steps = int(steps/denoise_strength) @@ -1913,101 +1528,27 @@ class WanVideoSampler: if steps != len(cfg): log.info(f"Received {len(cfg)} cfg values, but only {steps} steps. Setting step count to match.") steps = len(cfg) + else: + cfg = [cfg] * (steps +1) - - def get_scheduler(scheduler, steps, shift, device, sigmas=None): - timesteps = None - if 'unipc' in scheduler: - sample_scheduler = FlowUniPCMultistepScheduler(shift=shift) - if sigmas is None: - sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler)) - else: - sample_scheduler.sigmas = sigmas.to(device) - sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device) - sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps) + seed_g = torch.Generator(device=torch.device("cpu")) + seed_g.manual_seed(seed) - elif scheduler in ['euler/beta', 'euler']: - sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta')) - if flowedit_args: #seems to work better - timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift)) - else: - sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) - elif scheduler in ['euler/accvideo']: - if steps != 50: - raise Exception("Steps must be set to 50 for accvideo scheduler, 10 actual steps are used") - sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta')) - sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) - start_latent_list = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50] - sample_scheduler.sigmas = sample_scheduler.sigmas[start_latent_list] - steps = len(start_latent_list) - 1 - sample_scheduler.timesteps = timesteps = sample_scheduler.timesteps[start_latent_list[:steps]] - elif 'dpm++' in scheduler: - if 'sde' in scheduler: - algorithm_type = "sde-dpmsolver++" - else: - algorithm_type = "dpmsolver++" - sample_scheduler = FlowDPMSolverMultistepScheduler(shift=shift, algorithm_type=algorithm_type) - if sigmas is None: - sample_scheduler.set_timesteps(steps, device=device, use_beta_sigmas=('beta' in scheduler)) - else: - sample_scheduler.sigmas = sigmas.to(device) - sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device) - sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps) - elif scheduler == 'deis': - sample_scheduler = DEISMultistepScheduler(use_flow_sigmas=True, prediction_type="flow_prediction", flow_shift=shift) - sample_scheduler.set_timesteps(steps, device=device) - sample_scheduler.sigmas[-1] = 1e-6 - elif 'lcm' in scheduler: - sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta')) - sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) - elif 'flowmatch_causvid' in scheduler: - if transformer.dim == 5120: - denoising_list = [999, 934, 862, 756, 603, 410, 250, 140, 74] - else: - if steps != 4: - raise ValueError("CausVid 1.3B schedule is only for 4 steps") - denoising_list = [1000, 750, 500, 250] - sample_scheduler = FlowMatchScheduler(num_inference_steps=steps, shift=shift, sigma_min=0, extra_one_step=True) - sample_scheduler.timesteps = torch.tensor(denoising_list)[:steps].to(device) - sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)]) - elif 'flowmatch_distill' in scheduler: - sample_scheduler = FlowMatchScheduler( - shift=shift, sigma_min=0.0, extra_one_step=True - ) - sample_scheduler.set_timesteps(1000, training=True) - - denoising_step_list = torch.tensor([999, 750, 500, 250] , dtype=torch.long) - temp_timesteps = torch.cat((sample_scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32))) - denoising_step_list = temp_timesteps[1000 - denoising_step_list] - #print("denoising_step_list: ", denoising_step_list) - - if steps != 4: - raise ValueError("This scheduler is only for 4 steps") - - sample_scheduler.timesteps = denoising_step_list[:steps].clone().detach().to(device) - sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)]) - elif 'flowmatch_pusa' in scheduler: - sample_scheduler = FlowMatchSchedulerPusa( - shift=shift, sigma_min=0.0, extra_one_step=True - ) - sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, shift=shift) - - return sample_scheduler, timesteps - + # Scheduler if scheduler != "multitalk": - sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, sigmas=sigmas) - if timesteps is None: - timesteps = sample_scheduler.timesteps - log.info(f"timesteps: {timesteps}") + sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) else: timesteps = torch.tensor([1000, 750, 500, 250], device=device) + + scheduler_step_args = {"generator": seed_g} + step_sig = inspect.signature(sample_scheduler.step) + for arg in list(scheduler_step_args.keys()): + if arg not in step_sig.parameters: + scheduler_step_args.pop(arg) if denoise_strength < 1.0: steps = int(steps * denoise_strength) timesteps = timesteps[-(steps + 1):] - - seed_g = torch.Generator(device=torch.device("cpu")) - seed_g.manual_seed(seed) control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = None vace_data = vace_context = vace_scale = None @@ -2197,6 +1738,7 @@ class WanVideoSampler: if saved_generator_state is not None: seed_g.set_state(saved_generator_state) + # UniAnimate if unianimate_poses is not None: transformer.dwpose_embedding.to(device, model["dtype"]) dwpose_data = unianimate_poses["pose"].to(device, model["dtype"]) @@ -2230,6 +1772,7 @@ class WanVideoSampler: "end_percent": unianimate_poses["end_percent"] } + # FantasyTalking audio_proj = multitalk_audio_embedding = None audio_scale = 1.0 if fantasytalking_embeds is not None: @@ -2261,7 +1804,7 @@ class WanVideoSampler: shapes = [tuple(e.shape) for e in multitalk_audio_embedding] log.info(f"Multitalk audio features shapes (per speaker): {shapes}") - + # MiniMax Remover minimax_latents = minimax_mask_latents = None minimax_latents = image_embeds.get("minimax_latents", None) minimax_mask_latents = image_embeds.get("minimax_mask_latents", None) @@ -2271,58 +1814,9 @@ class WanVideoSampler: minimax_latents = minimax_latents.to(device, dtype) minimax_mask_latents = minimax_mask_latents.to(device, dtype) + # Context windows is_looped = False - if context_options is not None: - 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) - - 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)) - - # 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 - + if context_options is not None: context_schedule = context_options["context_schedule"] context_frames = (context_options["context_frames"] - 1) // 4 + 1 context_stride = context_options["context_stride"] // 4 @@ -2331,8 +1825,6 @@ class WanVideoSampler: if context_vae is not None: context_vae.to(device) - self.window_tracker = WindowTracker(verbose=context_options["verbose"]) - # Get total number of prompts num_prompts = len(text_embeds["prompt_embeds"]) log.info(f"Number of prompts: {num_prompts}") @@ -2364,9 +1856,11 @@ class WanVideoSampler: noise[:, place_idx:place_idx + delta, :, :] = noise[:, list_idx, :, :] log.info(f"Context schedule enabled: {context_frames} frames, {context_stride} stride, {context_overlap} overlap") - from .context import get_context_scheduler + from .context_windows.context import get_context_scheduler, create_window_mask, WindowTracker + self.window_tracker = WindowTracker(verbose=context_options["verbose"]) context = get_context_scheduler(context_schedule) + # vid2vid if samples is not None: input_samples = samples["samples"].squeeze(0).to(noise) if input_samples.shape[1] != noise.shape[1]: @@ -2380,7 +1874,8 @@ class WanVideoSampler: if mask is not None: if mask.shape[2] != noise.shape[1]: mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2) - + + # extra latents (Pusa) if (extra_latents := image_embeds.get("extra_latents", None)) is not None: encoded_image_latents = extra_latents["samples"].squeeze(0).to(noise) if (empty_latent_indices := extra_latents.get("empty_latent_indices", None)) is not None and len(empty_latent_indices) > 0: @@ -2394,72 +1889,6 @@ class WanVideoSampler: latent = noise.to(device) - freqs = None - transformer.rope_embedder.k = None - transformer.rope_embedder.num_frames = None - if "comfy" in rope_function: - transformer.rope_embedder.k = riflex_freq_index - transformer.rope_embedder.num_frames = latent_video_length - else: - d = transformer.dim // transformer.num_heads - freqs = torch.cat([ - rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index), - rope_params(1024, 2 * (d // 6)), - rope_params(1024, 2 * (d // 6)) - ], - dim=1) - transformer.rope_func = rope_function - for block in transformer.blocks: - block.rope_func = rope_function - if transformer.vace_layers is not None: - for block in transformer.vace_blocks: - block.rope_func = rope_function - - if not isinstance(cfg, list): - cfg = [cfg] * (steps +1) - - log.info(f"Seq len: {seq_len}") - - pbar = ProgressBar(steps) - - if args.preview_method in [LatentPreviewMethod.Auto, LatentPreviewMethod.Latent2RGB]: #default for latent2rgb - from latent_preview import prepare_callback - else: - from .latent_preview import prepare_callback #custom for tiny VAE previews - callback = prepare_callback(patcher, steps) - - #blockswap init - if transformer_options is not None: - block_swap_args = transformer_options.get("block_swap_args", None) - - if block_swap_args is not None: - transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True) - for name, param in transformer.named_parameters(): - if "block" not in name: - param.data = param.data.to(device) - if "control_adapter" in name: - param.data = param.data.to(device) - elif block_swap_args["offload_txt_emb"] and "txt_emb" in name: - param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking) - elif block_swap_args["offload_img_emb"] and "img_emb" in name: - param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking) - - transformer.block_swap( - block_swap_args["blocks_to_swap"] - 1 , - block_swap_args["offload_txt_emb"], - block_swap_args["offload_img_emb"], - vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None), - ) - - elif model["auto_cpu_offload"]: - for module in transformer.modules(): - if hasattr(module, "offload"): - module.offload() - if hasattr(module, "onload"): - module.onload() - elif model["manual_offloading"]: - transformer.to(device) - #controlnet controlnet_latents = controlnet = None if transformer_options is not None: @@ -2485,7 +1914,7 @@ class WanVideoSampler: "end": uni3c_embeds["end"], } - #feta + # Enhance-a-video (feta) if feta_args is not None and latent_video_length > 1: set_enhance_weight(feta_args["weight"]) feta_start_percent = feta_args["start_percent"] @@ -2499,37 +1928,76 @@ class WanVideoSampler: feta_args = None enhance_enabled = False - # Initialize Cache if enabled - transformer.enable_teacache = transformer.enable_magcache = False - if teacache_args is not None: #for backward compatibility on old workflows - cache_args = teacache_args - if cache_args is not None: - transformer.cache_device = cache_args["cache_device"] - if cache_args["cache_type"] == "TeaCache": - log.info(f"TeaCache: Using cache device: {transformer.cache_device}") - transformer.teacache_state.clear_all() - transformer.enable_teacache = True - transformer.rel_l1_thresh = cache_args["rel_l1_thresh"] - transformer.teacache_start_step = cache_args["start_step"] - transformer.teacache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] - transformer.teacache_use_coefficients = cache_args["use_coefficients"] - transformer.teacache_mode = cache_args["mode"] - elif cache_args["cache_type"] == "MagCache": - log.info(f"MagCache: Using cache device: {transformer.cache_device}") - transformer.magcache_state.clear_all() - transformer.enable_magcache = True - transformer.magcache_start_step = cache_args["start_step"] - transformer.magcache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] - transformer.magcache_thresh = cache_args["magcache_thresh"] - transformer.magcache_K = cache_args["magcache_K"] - elif cache_args["cache_type"] == "EasyCache": - log.info(f"EasyCache: Using cache device: {transformer.cache_device}") - transformer.easycache_state.clear_all() - transformer.enable_easycache = True - transformer.easycache_start_step = cache_args["start_step"] - transformer.easycache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] - transformer.easycache_thresh = cache_args["easycache_thresh"] + #region transformer settings + #rope + freqs = None + transformer.rope_embedder.k = None + transformer.rope_embedder.num_frames = None + if "comfy" in rope_function: + transformer.rope_embedder.k = riflex_freq_index + transformer.rope_embedder.num_frames = latent_video_length + else: + d = transformer.dim // transformer.num_heads + freqs = torch.cat([ + rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index), + rope_params(1024, 2 * (d // 6)), + rope_params(1024, 2 * (d // 6)) + ], + dim=1) + transformer.rope_func = rope_function + for block in transformer.blocks: + block.rope_func = rope_function + if transformer.vace_layers is not None: + for block in transformer.vace_blocks: + block.rope_func = rope_function + #blockswap init + if transformer_options is not None: + block_swap_args = transformer_options.get("block_swap_args", None) + + if block_swap_args is not None: + transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True) + for name, param in transformer.named_parameters(): + if "block" not in name: + param.data = param.data.to(device) + if "control_adapter" in name: + param.data = param.data.to(device) + elif block_swap_args["offload_txt_emb"] and "txt_emb" in name: + param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking) + elif block_swap_args["offload_img_emb"] and "img_emb" in name: + param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking) + + transformer.block_swap( + block_swap_args["blocks_to_swap"] - 1 , + block_swap_args["offload_txt_emb"], + block_swap_args["offload_img_emb"], + vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None), + ) + elif model["auto_cpu_offload"]: + for module in transformer.modules(): + if hasattr(module, "offload"): + module.offload() + if hasattr(module, "onload"): + module.onload() + elif model["manual_offloading"]: + transformer.to(device) + + # Initialize Cache if enabled + transformer.enable_teacache = transformer.enable_magcache = transformer.enable_easycache = False + cache_args = teacache_args if teacache_args is not None else cache_args #for backward compatibility on old workflows + if cache_args is not None: + from .cache_methods.cache_methods import set_transformer_cache_method + transformer = set_transformer_cache_method(transformer, timesteps, cache_args) + + # Initialize cache state + self.cache_state = [None, None] + if phantom_latents is not None: + log.info(f"Phantom latents shape: {phantom_latents.shape}") + self.cache_state = [None, None, None] + self.cache_state_source = [None, None] + self.cache_states_context = [] + + # Skip layer guidance (SLG) if slg_args is not None: assert batched_cfg is not None, "Batched cfg is not supported with SLG" transformer.slg_blocks = slg_args["blocks"] @@ -2570,13 +2038,8 @@ class WanVideoSampler: log.info(f"Radial attention mode enabled. dense_attention_mode: {dense_attention_mode}, dense_timesteps: {dense_timesteps}, dense_blocks: {dense_blocks}, decay_factor: {decay_factor}") - self.cache_state = [None, None] - if phantom_latents is not None: - log.info(f"Phantom latents shape: {phantom_latents.shape}") - self.cache_state = [None, None, None] - self.cache_state_source = [None, None] - self.cache_states_context = [] + # FlowEdit setup if flowedit_args is not None: source_embeds = flowedit_args["source_embeds"] source_image_embeds = flowedit_args.get("source_image_embeds", image_embeds) @@ -2612,6 +2075,7 @@ class WanVideoSampler: drift_timesteps = torch.cat([drift_timesteps, torch.tensor([0]).to(drift_timesteps.device)]).to(drift_timesteps.device) timesteps[-drift_steps:] = drift_timesteps[-drift_steps:] + # Experimental args use_cfg_zero_star = use_fresca = False if experimental_args is not None: video_attention_split_steps = experimental_args.get("video_attention_split_steps", []) @@ -2916,6 +2380,16 @@ class WanVideoSampler: return noise_pred, [cache_state_cond, cache_state_uncond] + + log.info(f"Seq len: {seq_len}") + + pbar = ProgressBar(steps) + + if args.preview_method in [LatentPreviewMethod.Auto, LatentPreviewMethod.Latent2RGB]: #default for latent2rgb + from latent_preview import prepare_callback + else: + from .latent_preview import prepare_callback #custom for tiny VAE previews + callback = prepare_callback(patcher, steps) log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*8}x{latent.shape[2]*8} with {steps} steps") @@ -2957,9 +2431,7 @@ class WanVideoSampler: # FreeInit noise reinitialization (after first iteration) if freeinit_args is not None and iter_idx > 0: # restart scheduler for each iteration - sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, sigmas=sigmas) - if timesteps is None: - timesteps = sample_scheduler.timesteps + sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) # Diffuse current latent to t=999 diffuse_timesteps = torch.full((noise.shape[0],), 999, device=device, dtype=torch.long) @@ -3333,9 +2805,7 @@ class WanVideoSampler: timesteps = [torch.tensor([t], device=device) for t in timesteps] timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps] else: - sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, sigmas=sigmas) - if timesteps is None: - timesteps = sample_scheduler.timesteps + sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) transformed_timesteps = [] for t in timesteps: @@ -3414,17 +2884,12 @@ class WanVideoSampler: latent = latent + noise_pred * dt[:, None, None, None] else: latent = latent.to(intermediate_device) - step_args = { - "generator": seed_g, - } - if isinstance(sample_scheduler, DEISMultistepScheduler) or isinstance(sample_scheduler, FlowMatchScheduler): - step_args.pop("generator", None) + temp_x0 = sample_scheduler.step( noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0), - #return_dict=False, - **step_args)[0] + **scheduler_step_args)[0] latent = temp_x0.squeeze(0) # injecting motion frames @@ -3532,17 +2997,11 @@ class WanVideoSampler: if flowedit_args is None: latent = latent.to(intermediate_device) - step_args = { - "generator": seed_g, - } - if isinstance(sample_scheduler, DEISMultistepScheduler) or isinstance(sample_scheduler, FlowMatchScheduler): - step_args.pop("generator", None) temp_x0 = sample_scheduler.step( noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else noise_pred.unsqueeze(0), timestep, latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else latent.unsqueeze(0), - #return_dict=False, - **step_args)[0] + **scheduler_step_args)[0] latent = temp_x0.squeeze(0) x0 = latent.to(device) @@ -3572,25 +3031,7 @@ class WanVideoSampler: x0 = x0[:,:-phantom_latents.shape[1]] if cache_args is not None: - cache_type = cache_args["cache_type"] - states = ( - transformer.teacache_state.states if cache_type == "TeaCache" else - transformer.magcache_state.states if cache_type == "MagCache" else - transformer.easycache_state.states if cache_type == "EasyCache" else - None - ) - state_names = { - 0: "conditional", - 1: "unconditional" - } - for pred_id, state in states.items(): - name = state_names.get(pred_id, f"prediction_{pred_id}") - if 'skipped_steps' in state: - log.info(f"{cache_type} skipped: {len(state['skipped_steps'])} {name} steps: {state['skipped_steps']}") - transformer.teacache_state.clear_all() - transformer.magcache_state.clear_all() - transformer.easycache_state.clear_all() - del states + cache_report(transformer, cache_args) if force_offload: if model["manual_offloading"]: @@ -3614,29 +3055,6 @@ class WanVideoSampler: "drop_last": drop_last, "generator_state": seed_g.get_state(), }, ) - -class WindowTracker: - def __init__(self, verbose=False): - self.window_map = {} # Maps frame sequence to persistent ID - self.next_id = 0 - self.cache_states = {} # Maps persistent ID to teacache state - self.verbose = verbose - - def get_window_id(self, frames): - key = tuple(sorted(frames)) # Order-independent frame sequence - if key not in self.window_map: - self.window_map[key] = self.next_id - if self.verbose: - log.info(f"New window pattern {key} -> ID {self.next_id}") - self.next_id += 1 - return self.window_map[key] - - def get_teacache(self, window_id, base_state): - if window_id not in self.cache_states: - if self.verbose: - log.info(f"Initializing persistent teacache for window {window_id}") - self.cache_states[window_id] = base_state.copy() - return self.cache_states[window_id] #region VideoDecode class WanVideoDecode: @@ -3676,8 +3094,6 @@ class WanVideoDecode: CATEGORY = "WanVideoWrapper" def decode(self, vae, samples, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, normalization="default"): - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() mm.soft_empty_cache() video = samples.get("video", None) if video is not None: @@ -3769,9 +3185,6 @@ class WanVideoEncode: CATEGORY = "WanVideoWrapper" def encode(self, vae, image, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, noise_aug_strength=0.0, latent_strength=1.0, mask=None): - device = mm.get_torch_device() - offload_device = mm.unet_offload_device() - vae.to(device) image = image.clone() @@ -3866,23 +3279,16 @@ NODE_CLASS_MAPPINGS = { "WanVideoEmptyEmbeds": WanVideoEmptyEmbeds, "WanVideoEnhanceAVideo": WanVideoEnhanceAVideo, "WanVideoContextOptions": WanVideoContextOptions, - "WanVideoTeaCache": WanVideoTeaCache, - "WanVideoMagCache": WanVideoMagCache, - "WanVideoEasyCache": WanVideoEasyCache, "WanVideoVRAMManagement": WanVideoVRAMManagement, "WanVideoTextEmbedBridge": WanVideoTextEmbedBridge, "WanVideoFlowEdit": WanVideoFlowEdit, "WanVideoControlEmbeds": WanVideoControlEmbeds, "WanVideoSLG": WanVideoSLG, "WanVideoLoopArgs": WanVideoLoopArgs, - "WanVideoImageResizeToClosest": WanVideoImageResizeToClosest, "WanVideoSetBlockSwap": WanVideoSetBlockSwap, "WanVideoExperimentalArgs": WanVideoExperimentalArgs, "WanVideoVACEEncode": WanVideoVACEEncode, - "ExtractStartFramesForContinuations": ExtractStartFramesForContinuations, - "WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame, "WanVideoPhantomEmbeds": WanVideoPhantomEmbeds, - "CreateCFGScheduleFloatList": CreateCFGScheduleFloatList, "WanVideoRealisDanceLatents": WanVideoRealisDanceLatents, "WanVideoApplyNAG": WanVideoApplyNAG, "WanVideoMiniMaxRemoverEmbeds": WanVideoMiniMaxRemoverEmbeds, @@ -3906,23 +3312,16 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoEmptyEmbeds": "WanVideo Empty Embeds", "WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video", "WanVideoContextOptions": "WanVideo Context Options", - "WanVideoTeaCache": "WanVideo TeaCache", - "WanVideoMagCache": "WanVideo MagCache", - "WanVideoEasyCache": "WanVideo EasyCache", "WanVideoVRAMManagement": "WanVideo VRAM Management", "WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge", "WanVideoFlowEdit": "WanVideo FlowEdit", "WanVideoControlEmbeds": "WanVideo Control Embeds", "WanVideoSLG": "WanVideo SLG", "WanVideoLoopArgs": "WanVideo Loop Args", - "WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest", "WanVideoSetBlockSwap": "WanVideo Set BlockSwap", "WanVideoExperimentalArgs": "WanVideo Experimental Args", "WanVideoVACEEncode": "WanVideo VACE Encode", - "ExtractStartFramesForContinuations": "Extract Start Frames For Continuations", - "WanVideoVACEStartToEndFrame": "WanVideo VACE Start To End Frame", "WanVideoPhantomEmbeds": "WanVideo Phantom Embeds", - "CreateCFGScheduleFloatList": "WanVideo CFG Schedule Float List", "WanVideoRealisDanceLatents": "WanVideo RealisDance Latents", "WanVideoApplyNAG": "WanVideo Apply NAG", "WanVideoMiniMaxRemoverEmbeds": "WanVideo MiniMax Remover Embeds", diff --git a/nodes_utility.py b/nodes_utility.py new file mode 100644 index 0000000..5a12804 --- /dev/null +++ b/nodes_utility.py @@ -0,0 +1,242 @@ +import torch +import numpy as np +from comfy.utils import common_upscale + +VAE_STRIDE = (4, 8, 8) +PATCH_SIZE = (1, 2, 2) + +class WanVideoImageResizeToClosest: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image": ("IMAGE", {"tooltip": "Image to resize"}), + "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"],), + }, + } + + RETURN_TYPES = ("IMAGE", "INT", "INT", ) + RETURN_NAMES = ("image","width","height",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Resizes image to the closest supported resolution based on aspect ratio and max pixels, according to the original code" + + def process(self, image, generation_width, generation_height, aspect_ratio_preservation ): + + H, W = image.shape[1], image.shape[2] + max_area = generation_width * generation_height + + crop = "disabled" + + if aspect_ratio_preservation == "keep_input": + aspect_ratio = H / W + elif aspect_ratio_preservation == "stretch_to_new" or aspect_ratio_preservation == "crop_to_new": + aspect_ratio = generation_height / generation_width + if aspect_ratio_preservation == "crop_to_new": + crop = "center" + + lat_h = round( + np.sqrt(max_area * aspect_ratio) // VAE_STRIDE[1] // + PATCH_SIZE[1] * PATCH_SIZE[1]) + lat_w = round( + np.sqrt(max_area / aspect_ratio) // VAE_STRIDE[2] // + PATCH_SIZE[2] * PATCH_SIZE[2]) + h = lat_h * VAE_STRIDE[1] + w = lat_w * VAE_STRIDE[2] + + resized_image = common_upscale(image.movedim(-1, 1), w, h, "lanczos", crop).movedim(1, -1) + + return (resized_image, w, h) + +class ExtractStartFramesForContinuations: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "input_video_frames": ("IMAGE", {"tooltip": "Input video frames to extract the start frames from."}), + "num_frames": ("INT", {"default": 10, "min": 1, "max": 1024, "step": 1, "tooltip": "Number of frames to get from the start of the video."}), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("start_frames",) + FUNCTION = "get_start_frames" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Extracts the first N frames from a video sequence for continuations." + + def get_start_frames(self, input_video_frames, num_frames): + if input_video_frames is None or input_video_frames.shape[0] == 0: + log.warning("Input video frames are empty. Returning an empty tensor.") + if input_video_frames is not None: + return (torch.empty((0,) + input_video_frames.shape[1:], dtype=input_video_frames.dtype),) + else: + # Return a tensor with 4 dimensions, as expected for an IMAGE type. + return (torch.empty((0, 64, 64, 3), dtype=torch.float32),) + + total_frames = input_video_frames.shape[0] + num_to_get = min(num_frames, total_frames) + + if num_to_get < num_frames: + log.warning(f"Requested {num_frames} frames, but input video only has {total_frames} frames. Returning first {num_to_get} frames.") + + start_frames = input_video_frames[:num_to_get] + + return (start_frames.cpu().float(),) + +class WanVideoVACEStartToEndFrame: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), + "empty_frame_level": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "White level of empty frame to use"}), + }, + "optional": { + "start_image": ("IMAGE",), + "end_image": ("IMAGE",), + "control_images": ("IMAGE",), + "inpaint_mask": ("MASK", {"tooltip": "Inpaint mask to use for the empty frames"}), + "start_index": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Index to start from"}), + "end_index": ("INT", {"default": -1, "min": -10000, "max": 10000, "step": 1, "tooltip": "Index to end at"}), + }, + } + + RETURN_TYPES = ("IMAGE", "MASK", ) + RETURN_NAMES = ("images", "masks",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Helper node to create start/end frame batch and masks for VACE" + + def process(self, num_frames, empty_frame_level, start_image=None, end_image=None, control_images=None, inpaint_mask=None, start_index=0, end_index=-1): + + B, H, W, C = start_image.shape if start_image is not None else end_image.shape + device = start_image.device if start_image is not None else end_image.device + + # Convert negative end_index to positive + if end_index < 0: + end_index = num_frames + end_index + + # Create output batch with empty frames + out_batch = torch.ones((num_frames, H, W, 3), device=device) * empty_frame_level + + # Create mask tensor with proper dimensions + masks = torch.ones((num_frames, H, W), device=device) + + # Pre-process all images at once to avoid redundant work + if end_image is not None and (end_image.shape[1] != H or end_image.shape[2] != W): + end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1) + + if control_images is not None and (control_images.shape[1] != H or control_images.shape[2] != W): + control_images = common_upscale(control_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1) + + # Place start image at start_index + if start_image is not None: + frames_to_copy = min(start_image.shape[0], num_frames - start_index) + if frames_to_copy > 0: + out_batch[start_index:start_index + frames_to_copy] = start_image[:frames_to_copy] + masks[start_index:start_index + frames_to_copy] = 0 + + # Place end image at end_index + if end_image is not None: + # Calculate where to start placing end images + end_start = end_index - end_image.shape[0] + 1 + if end_start < 0: # Handle case where end images won't all fit + end_image = end_image[abs(end_start):] + end_start = 0 + + frames_to_copy = min(end_image.shape[0], num_frames - end_start) + if frames_to_copy > 0: + out_batch[end_start:end_start + frames_to_copy] = end_image[:frames_to_copy] + masks[end_start:end_start + frames_to_copy] = 0 + + # Apply control images to remaining frames that don't have start or end images + if control_images is not None: + # Create a mask of frames that are still empty (mask == 1) + empty_frames = masks.sum(dim=(1, 2)) > 0.5 * H * W + + if empty_frames.any(): + # Only apply control images where they exist + control_length = control_images.shape[0] + for frame_idx in range(num_frames): + if empty_frames[frame_idx] and frame_idx < control_length: + out_batch[frame_idx] = control_images[frame_idx] + + # Apply inpaint mask if provided + if inpaint_mask is not None: + inpaint_mask = common_upscale(inpaint_mask.unsqueeze(1), W, H, "nearest-exact", "disabled").squeeze(1).to(device) + + # Handle different mask lengths efficiently + if inpaint_mask.shape[0] > num_frames: + inpaint_mask = inpaint_mask[:num_frames] + elif inpaint_mask.shape[0] < num_frames: + repeat_factor = (num_frames + inpaint_mask.shape[0] - 1) // inpaint_mask.shape[0] # Ceiling division + inpaint_mask = inpaint_mask.repeat(repeat_factor, 1, 1)[:num_frames] + + # Apply mask in one operation + masks = inpaint_mask * masks + + return (out_batch.cpu().float(), masks.cpu().float()) + + +class CreateCFGScheduleFloatList: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "steps": ("INT", {"default": 30, "min": 2, "max": 1000, "step": 1, "tooltip": "Number of steps to schedule cfg for"} ), + "cfg_scale_start": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 30.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}), + "cfg_scale_end": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 30.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}), + "interpolation": (["linear", "ease_in", "ease_out"], {"default": "linear", "tooltip": "Interpolation method to use for the cfg scale"}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "Start percent of the steps to apply cfg"}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "End percent of the steps to apply cfg"}), + } + } + + RETURN_TYPES = ("FLOAT", ) + RETURN_NAMES = ("float_list",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Helper node to generate a list of floats that can be used to schedule cfg scale for the steps, outside the set range cfg is set to 1.0" + + def process(self, steps, cfg_scale_start, cfg_scale_end, interpolation, start_percent, end_percent): + + # Create a list of floats for the cfg schedule + cfg_list = [1.0] * steps + start_idx = min(int(steps * start_percent), steps - 1) + end_idx = min(int(steps * end_percent), steps - 1) + + for i in range(start_idx, end_idx + 1): + if i >= steps: + break + + if end_idx == start_idx: + t = 0 + else: + t = (i - start_idx) / (end_idx - start_idx) + + if interpolation == "linear": + factor = t + elif interpolation == "ease_in": + factor = t * t + elif interpolation == "ease_out": + factor = t * (2 - t) + + cfg_list[i] = round(cfg_scale_start + factor * (cfg_scale_end - cfg_scale_start), 2) + + # If start_percent > 0, always include the first step + if start_percent > 0: + cfg_list[0] = 1.0 + + return (cfg_list,) + +NODE_CLASS_MAPPINGS = { + "WanVideoImageResizeToClosest": WanVideoImageResizeToClosest, + "WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame, + "ExtractStartFramesForContinuations": ExtractStartFramesForContinuations, + "CreateCFGScheduleFloatList": CreateCFGScheduleFloatList + } +NODE_DISPLAY_NAME_MAPPINGS = { + "WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest", + "WanVideoVACEStartToEndFrame": "WanVideo VACE Start To End Frame", + "ExtractStartFramesForContinuations": "Extract Start Frames For Continuations", + "CreateCFGScheduleFloatList": "Create CFG Schedule Float List" + } \ No newline at end of file diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 31664ea..8c4d6ff 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -28,6 +28,7 @@ from tqdm import tqdm import gc import comfy.model_management as mm from ...utils import log, get_module_memory_mb +from ...cache_methods.cache_methods import TeaCacheState, MagCacheState, EasyCacheState, relative_l1_distance from ...multitalk.multitalk import get_attn_map_with_target @@ -1867,110 +1868,3 @@ class WanModel(ModelMixin, ConfigMixin): u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)]) out.append(u) return out - -class TeaCacheState: - def __init__(self, cache_device='cpu'): - self.cache_device = cache_device - self.states = {} - self._next_pred_id = 0 - - def new_prediction(self, cache_device='cpu'): - """Create new prediction state and return its ID""" - self.cache_device = cache_device - pred_id = self._next_pred_id - self._next_pred_id += 1 - self.states[pred_id] = { - 'previous_residual': None, - 'accumulated_rel_l1_distance': 0, - 'previous_modulated_input': None, - 'skipped_steps': [], - } - return pred_id - - def update(self, pred_id, **kwargs): - """Update state for specific prediction""" - if pred_id not in self.states: - return None - for key, value in kwargs.items(): - self.states[pred_id][key] = value - - def get(self, pred_id): - return self.states.get(pred_id, {}) - - def clear_all(self): - self.states = {} - self._next_pred_id = 0 - -class MagCacheState: - def __init__(self, cache_device='cpu'): - self.cache_device = cache_device - self.states = {} - self._next_pred_id = 0 - - def new_prediction(self, cache_device='cpu'): - """Create new prediction state and return its ID""" - self.cache_device = cache_device - pred_id = self._next_pred_id - self._next_pred_id += 1 - self.states[pred_id] = { - 'residual_cache': None, - 'accumulated_ratio': 1.0, - 'accumulated_steps': 0, - 'accumulated_err': 0, - 'skipped_steps': [], - } - return pred_id - - def update(self, pred_id, **kwargs): - """Update state for specific prediction""" - if pred_id not in self.states: - return None - for key, value in kwargs.items(): - self.states[pred_id][key] = value - - def get(self, pred_id): - return self.states.get(pred_id, {}) - - def clear_all(self): - self.states = {} - self._next_pred_id = 0 - -class EasyCacheState: - def __init__(self, cache_device='cpu'): - self.cache_device = cache_device - self.states = {} - self._next_pred_id = 0 - - def new_prediction(self, cache_device='cpu'): - """Create a new prediction state and return its ID.""" - self.cache_device = cache_device - pred_id = self._next_pred_id - self._next_pred_id += 1 - self.states[pred_id] = { - 'previous_raw_input': None, - 'previous_raw_output': None, - 'cache': None, - 'accumulated_error': 0.0, - 'skipped_steps': [], - } - return pred_id - - def update(self, pred_id, **kwargs): - """Update state for a specific prediction.""" - if pred_id not in self.states: - return None - for key, value in kwargs.items(): - self.states[pred_id][key] = value - - def get(self, pred_id): - return self.states.get(pred_id, {}) - - def clear_all(self): - self.states = {} - self._next_pred_id = 0 - -def relative_l1_distance(last_tensor, current_tensor): - l1_distance = torch.abs(last_tensor.to(current_tensor.device) - current_tensor).mean() - norm = torch.abs(last_tensor).mean() - relative_l1_distance = l1_distance / norm - return relative_l1_distance.to(torch.float32).to(current_tensor.device) diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index 3f5e8b5..1188b96 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -1,5 +1,90 @@ +import torch from .fm_solvers import (FlowDPMSolverMultistepScheduler, get_sampling_sigmas, retrieve_timesteps) from .fm_solvers_unipc import FlowUniPCMultistepScheduler from .basic_flowmatch import FlowMatchScheduler from .flowmatch_pusa import FlowMatchSchedulerPusa from .scheduling_flow_match_lcm import FlowMatchLCMScheduler +from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler + +from ...utils import log + +def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_args, denoise_strength, sigmas=None): + timesteps = None + if 'unipc' in scheduler: + sample_scheduler = FlowUniPCMultistepScheduler(shift=shift) + if sigmas is None: + sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler)) + else: + sample_scheduler.sigmas = sigmas.to(device) + sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device) + sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps) + + elif scheduler in ['euler/beta', 'euler']: + sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta')) + if flowedit_args: #seems to work better + timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift)) + else: + sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) + elif scheduler in ['euler/accvideo']: + if steps != 50: + raise Exception("Steps must be set to 50 for accvideo scheduler, 10 actual steps are used") + sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta')) + sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) + start_latent_list = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50] + sample_scheduler.sigmas = sample_scheduler.sigmas[start_latent_list] + steps = len(start_latent_list) - 1 + sample_scheduler.timesteps = timesteps = sample_scheduler.timesteps[start_latent_list[:steps]] + elif 'dpm++' in scheduler: + if 'sde' in scheduler: + algorithm_type = "sde-dpmsolver++" + else: + algorithm_type = "dpmsolver++" + sample_scheduler = FlowDPMSolverMultistepScheduler(shift=shift, algorithm_type=algorithm_type) + if sigmas is None: + sample_scheduler.set_timesteps(steps, device=device, use_beta_sigmas=('beta' in scheduler)) + else: + sample_scheduler.sigmas = sigmas.to(device) + sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device) + sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps) + elif scheduler == 'deis': + sample_scheduler = DEISMultistepScheduler(use_flow_sigmas=True, prediction_type="flow_prediction", flow_shift=shift) + sample_scheduler.set_timesteps(steps, device=device) + sample_scheduler.sigmas[-1] = 1e-6 + elif 'lcm' in scheduler: + sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta')) + sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) + elif 'flowmatch_causvid' in scheduler: + if transformer_dim == 5120: + denoising_list = [999, 934, 862, 756, 603, 410, 250, 140, 74] + else: + if steps != 4: + raise ValueError("CausVid 1.3B schedule is only for 4 steps") + denoising_list = [1000, 750, 500, 250] + sample_scheduler = FlowMatchScheduler(num_inference_steps=steps, shift=shift, sigma_min=0, extra_one_step=True) + sample_scheduler.timesteps = torch.tensor(denoising_list)[:steps].to(device) + sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)]) + elif 'flowmatch_distill' in scheduler: + sample_scheduler = FlowMatchScheduler( + shift=shift, sigma_min=0.0, extra_one_step=True + ) + sample_scheduler.set_timesteps(1000, training=True) + + denoising_step_list = torch.tensor([999, 750, 500, 250] , dtype=torch.long) + temp_timesteps = torch.cat((sample_scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32))) + denoising_step_list = temp_timesteps[1000 - denoising_step_list] + #print("denoising_step_list: ", denoising_step_list) + + if steps != 4: + raise ValueError("This scheduler is only for 4 steps") + + sample_scheduler.timesteps = denoising_step_list[:steps].clone().detach().to(device) + sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)]) + elif 'flowmatch_pusa' in scheduler: + sample_scheduler = FlowMatchSchedulerPusa( + shift=shift, sigma_min=0.0, extra_one_step=True + ) + sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, shift=shift) + if timesteps is None: + timesteps = sample_scheduler.timesteps + log.info(f"timesteps: {timesteps}") + return sample_scheduler, timesteps \ No newline at end of file