diff --git a/nodes.py b/nodes.py index 3df8cb9..14d2f1c 100644 --- a/nodes.py +++ b/nodes.py @@ -7,22 +7,20 @@ 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 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 .utils import log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, is_image_black, add_noise_to_reference_video, optimized_scale 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 from comfy import model_management as mm -from comfy.utils import load_torch_file, ProgressBar, common_upscale +from comfy.utils import ProgressBar, common_upscale from comfy.clip_vision import clip_preprocess, ClipVisionModel from comfy.cli_args import args, LatentPreviewMethod @@ -34,67 +32,6 @@ offload_device = mm.unet_offload_device() VAE_STRIDE = (4, 8, 8) PATCH_SIZE = (1, 2, 2) -def add_noise_to_reference_video(image, ratio=None): - sigma = torch.ones((image.shape[0],)).to(image.device, image.dtype) * ratio - image_noise = torch.randn_like(image) * sigma[:, None, None, None] - image_noise = torch.where(image==-1, torch.zeros_like(image), image_noise) - image = image + image_noise - return image - -def optimized_scale(positive_flat, negative_flat): - - # Calculate dot production - dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True) - - # Squared norm of uncondition - squared_norm = torch.sum(negative_flat ** 2, dim=1, keepdim=True) + 1e-8 - - # st_star = v_cond^T * v_uncond / ||v_uncond||^2 - st_star = dot_product / squared_norm - - return st_star - -class WanVideoBlockSwap: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 40, "step": 1, "tooltip": "Number of transformer blocks to swap, the 14B model has 40, while the 1.3B model has 30 blocks"}), - "offload_img_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload img_emb to offload_device"}), - "offload_txt_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload time_emb to offload_device"}), - }, - "optional": { - "use_non_blocking": ("BOOLEAN", {"default": True, "tooltip": "Use non-blocking memory transfer for offloading, reserves more RAM but is faster"}), - "vace_blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 15, "step": 1, "tooltip": "Number of VACE blocks to swap, the VACE model has 15 blocks"}), - }, - } - RETURN_TYPES = ("BLOCKSWAPARGS",) - RETURN_NAMES = ("block_swap_args",) - FUNCTION = "setargs" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Settings for block swapping, reduces VRAM use by swapping blocks to CPU memory" - - def setargs(self, **kwargs): - return (kwargs, ) - -class WanVideoVRAMManagement: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "offload_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Percentage of parameters to offload"}), - }, - } - RETURN_TYPES = ("VRAM_MANAGEMENTARGS",) - RETURN_NAMES = ("vram_management_args",) - FUNCTION = "setargs" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Alternative offloading method from DiffSynth-Studio, more aggressive in reducing memory use than block swapping, but can be slower" - - def setargs(self, **kwargs): - return (kwargs, ) - - class WanVideoEnhanceAVideo: @classmethod def INPUT_TYPES(s): @@ -180,190 +117,7 @@ class WanVideoSetRadialAttention: return (patcher,) -class WanVideoTorchCompileSettings: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "backend": (["inductor","cudagraphs"], {"default": "inductor"}), - "fullgraph": ("BOOLEAN", {"default": False, "tooltip": "Enable full graph mode"}), - "mode": (["default", "max-autotune", "max-autotune-no-cudagraphs", "reduce-overhead"], {"default": "default"}), - "dynamic": ("BOOLEAN", {"default": False, "tooltip": "Enable dynamic mode"}), - "dynamo_cache_size_limit": ("INT", {"default": 64, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.cache_size_limit"}), - "compile_transformer_blocks_only": ("BOOLEAN", {"default": True, "tooltip": "Compile only the transformer blocks, usually enough and can make compilation faster and less error prone"}), - }, - "optional": { - "dynamo_recompile_limit": ("INT", {"default": 128, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.recompile_limit"}), - }, - } - RETURN_TYPES = ("WANCOMPILEARGS",) - RETURN_NAMES = ("torch_compile_args",) - FUNCTION = "set_args" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch 2.5.0 is recommended" - - def set_args(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_transformer_blocks_only, dynamo_recompile_limit=128): - - compile_args = { - "backend": backend, - "fullgraph": fullgraph, - "mode": mode, - "dynamic": dynamic, - "dynamo_cache_size_limit": dynamo_cache_size_limit, - "dynamo_recompile_limit": dynamo_recompile_limit, - "compile_transformer_blocks_only": compile_transformer_blocks_only, - } - - return (compile_args, ) - #region TextEncode - -class LoadWanVideoT5TextEncoder: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model_name": (folder_paths.get_filename_list("text_encoders"), {"tooltip": "These models are loaded from 'ComfyUI/models/text_encoders'"}), - "precision": (["fp32", "bf16"], - {"default": "bf16"} - ), - }, - "optional": { - "load_device": (["main_device", "offload_device"], {"default": "offload_device"}), - "quantization": (['disabled', 'fp8_e4m3fn'], {"default": 'disabled', "tooltip": "optional quantization method"}), - } - } - - RETURN_TYPES = ("WANTEXTENCODER",) - RETURN_NAMES = ("wan_t5_model", ) - FUNCTION = "loadmodel" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Loads Wan text_encoder model from 'ComfyUI/models/LLM'" - - def loadmodel(self, model_name, precision, load_device="offload_device", quantization="disabled"): - text_encoder_load_device = device if load_device == "main_device" else offload_device - - tokenizer_path = os.path.join(script_directory, "configs", "T5_tokenizer") - - dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] - - model_path = folder_paths.get_full_path("text_encoders", model_name) - sd = load_torch_file(model_path, safe_load=True) - - if "token_embedding.weight" not in sd and "shared.weight" not in sd: - raise ValueError("Invalid T5 text encoder model, this node expects the 'umt5-xxl' model") - if "scaled_fp8" in sd: - raise ValueError("Invalid T5 text encoder model, fp8 scaled is not supported by this node") - - # Convert state dict keys from T5 format to the expected format - if "shared.weight" in sd: - log.info("Converting T5 text encoder model to the expected format...") - converted_sd = {} - - for key, value in sd.items(): - # Handle encoder block patterns - if key.startswith('encoder.block.'): - parts = key.split('.') - block_num = parts[2] - - # Self-attention components - if 'layer.0.SelfAttention' in key: - if key.endswith('.k.weight'): - new_key = f"blocks.{block_num}.attn.k.weight" - elif key.endswith('.o.weight'): - new_key = f"blocks.{block_num}.attn.o.weight" - elif key.endswith('.q.weight'): - new_key = f"blocks.{block_num}.attn.q.weight" - elif key.endswith('.v.weight'): - new_key = f"blocks.{block_num}.attn.v.weight" - elif 'relative_attention_bias' in key: - new_key = f"blocks.{block_num}.pos_embedding.embedding.weight" - else: - new_key = key - - # Layer norms - elif 'layer.0.layer_norm' in key: - new_key = f"blocks.{block_num}.norm1.weight" - elif 'layer.1.layer_norm' in key: - new_key = f"blocks.{block_num}.norm2.weight" - - # Feed-forward components - elif 'layer.1.DenseReluDense' in key: - if 'wi_0' in key: - new_key = f"blocks.{block_num}.ffn.gate.0.weight" - elif 'wi_1' in key: - new_key = f"blocks.{block_num}.ffn.fc1.weight" - elif 'wo' in key: - new_key = f"blocks.{block_num}.ffn.fc2.weight" - else: - new_key = key - else: - new_key = key - elif key == "shared.weight": - new_key = "token_embedding.weight" - elif key == "encoder.final_layer_norm.weight": - new_key = "norm.weight" - else: - new_key = key - converted_sd[new_key] = value - sd = converted_sd - - T5_text_encoder = T5EncoderModel( - text_len=512, - dtype=dtype, - device=text_encoder_load_device, - state_dict=sd, - tokenizer_path=tokenizer_path, - quantization=quantization - ) - text_encoder = { - "model": T5_text_encoder, - "dtype": dtype, - } - - return (text_encoder,) - -class LoadWanVideoClipTextEncoder: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model_name": (folder_paths.get_filename_list("clip_vision") + folder_paths.get_filename_list("text_encoders"), {"tooltip": "These models are loaded from 'ComfyUI/models/clip_vision'"}), - "precision": (["fp16", "fp32", "bf16"], - {"default": "fp16"} - ), - }, - "optional": { - "load_device": (["main_device", "offload_device"], {"default": "offload_device"}), - } - } - - RETURN_TYPES = ("CLIP_VISION",) - RETURN_NAMES = ("wan_clip_vision", ) - FUNCTION = "loadmodel" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Loads Wan clip_vision model from 'ComfyUI/models/clip_vision'" - - def loadmodel(self, model_name, precision, load_device="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] - - model_path = folder_paths.get_full_path("clip_vision", model_name) - # We also support legacy setups where the model is in the text_encoders folder - if model_path is None: - model_path = folder_paths.get_full_path("text_encoders", model_name) - sd = load_torch_file(model_path, safe_load=True) - if "log_scale" not in sd: - raise ValueError("Invalid CLIP model, this node expectes the 'open-clip-xlm-roberta-large-vit-huge-14' model") - - clip_model = CLIPModel(dtype=dtype, device=device, state_dict=sd) - clip_model.model.to(text_encoder_load_device) - del sd - - return (clip_model,) - - class WanVideoTextEncode: @classmethod def INPUT_TYPES(s): @@ -3268,18 +3022,13 @@ NODE_CLASS_MAPPINGS = { "WanVideoDecode": WanVideoDecode, "WanVideoTextEncode": WanVideoTextEncode, "WanVideoTextEncodeSingle": WanVideoTextEncodeSingle, - "LoadWanVideoT5TextEncoder": LoadWanVideoT5TextEncoder, "WanVideoImageClipEncode": WanVideoImageClipEncode,#deprecated "WanVideoClipVisionEncode": WanVideoClipVisionEncode, "WanVideoImageToVideoEncode": WanVideoImageToVideoEncode, - "LoadWanVideoClipTextEncoder": LoadWanVideoClipTextEncoder, "WanVideoEncode": WanVideoEncode, - "WanVideoBlockSwap": WanVideoBlockSwap, - "WanVideoTorchCompileSettings": WanVideoTorchCompileSettings, "WanVideoEmptyEmbeds": WanVideoEmptyEmbeds, "WanVideoEnhanceAVideo": WanVideoEnhanceAVideo, "WanVideoContextOptions": WanVideoContextOptions, - "WanVideoVRAMManagement": WanVideoVRAMManagement, "WanVideoTextEmbedBridge": WanVideoTextEmbedBridge, "WanVideoFlowEdit": WanVideoFlowEdit, "WanVideoControlEmbeds": WanVideoControlEmbeds, @@ -3293,7 +3042,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoApplyNAG": WanVideoApplyNAG, "WanVideoMiniMaxRemoverEmbeds": WanVideoMiniMaxRemoverEmbeds, "WanVideoFreeInitArgs": WanVideoFreeInitArgs, - "WanVideoSetRadialAttention": WanVideoSetRadialAttention + "WanVideoSetRadialAttention": WanVideoSetRadialAttention, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -3301,18 +3050,13 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoTextEncode": "WanVideo TextEncode", "WanVideoTextEncodeSingle": "WanVideo TextEncodeSingle", "WanVideoTextImageEncode": "WanVideo TextImageEncode (IP2V)", - "LoadWanVideoT5TextEncoder": "Load WanVideo T5 TextEncoder", "WanVideoImageClipEncode": "WanVideo ImageClip Encode (Deprecated)", "WanVideoClipVisionEncode": "WanVideo ClipVision Encode", "WanVideoImageToVideoEncode": "WanVideo ImageToVideo Encode", - "LoadWanVideoClipTextEncoder": "Load WanVideo Clip Encoder", "WanVideoEncode": "WanVideo Encode", - "WanVideoBlockSwap": "WanVideo BlockSwap", - "WanVideoTorchCompileSettings": "WanVideo Torch Compile Settings", "WanVideoEmptyEmbeds": "WanVideo Empty Embeds", "WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video", "WanVideoContextOptions": "WanVideo Context Options", - "WanVideoVRAMManagement": "WanVideo VRAM Management", "WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge", "WanVideoFlowEdit": "WanVideo FlowEdit", "WanVideoControlEmbeds": "WanVideo Control Embeds", diff --git a/nodes_model_loading.py b/nodes_model_loading.py index d5cb4aa..37e542d 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1,9 +1,12 @@ import torch -import gc +import os, gc from .utils import log, apply_lora import numpy as np from tqdm import tqdm + from .wanvideo.modules.model import WanModel +from .wanvideo.modules.t5 import T5EncoderModel +from .wanvideo.modules.clip import CLIPModel from accelerate import init_empty_weights from accelerate.utils import set_module_tensor_to_device @@ -14,6 +17,11 @@ from comfy.utils import load_torch_file, ProgressBar import comfy.model_base from comfy.sd import load_lora_for_models +script_directory = os.path.dirname(os.path.abspath(__file__)) + +device = mm.get_torch_device() +offload_device = mm.unet_offload_device() + try: from server import PromptServer except: @@ -242,6 +250,83 @@ def standardize_lora_key_format(lora_sd): new_sd[k] = v return new_sd +class WanVideoBlockSwap: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 40, "step": 1, "tooltip": "Number of transformer blocks to swap, the 14B model has 40, while the 1.3B model has 30 blocks"}), + "offload_img_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload img_emb to offload_device"}), + "offload_txt_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload time_emb to offload_device"}), + }, + "optional": { + "use_non_blocking": ("BOOLEAN", {"default": True, "tooltip": "Use non-blocking memory transfer for offloading, reserves more RAM but is faster"}), + "vace_blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 15, "step": 1, "tooltip": "Number of VACE blocks to swap, the VACE model has 15 blocks"}), + }, + } + RETURN_TYPES = ("BLOCKSWAPARGS",) + RETURN_NAMES = ("block_swap_args",) + FUNCTION = "setargs" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Settings for block swapping, reduces VRAM use by swapping blocks to CPU memory" + + def setargs(self, **kwargs): + return (kwargs, ) + +class WanVideoVRAMManagement: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "offload_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Percentage of parameters to offload"}), + }, + } + RETURN_TYPES = ("VRAM_MANAGEMENTARGS",) + RETURN_NAMES = ("vram_management_args",) + FUNCTION = "setargs" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Alternative offloading method from DiffSynth-Studio, more aggressive in reducing memory use than block swapping, but can be slower" + + def setargs(self, **kwargs): + return (kwargs, ) + +class WanVideoTorchCompileSettings: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "backend": (["inductor","cudagraphs"], {"default": "inductor"}), + "fullgraph": ("BOOLEAN", {"default": False, "tooltip": "Enable full graph mode"}), + "mode": (["default", "max-autotune", "max-autotune-no-cudagraphs", "reduce-overhead"], {"default": "default"}), + "dynamic": ("BOOLEAN", {"default": False, "tooltip": "Enable dynamic mode"}), + "dynamo_cache_size_limit": ("INT", {"default": 64, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.cache_size_limit"}), + "compile_transformer_blocks_only": ("BOOLEAN", {"default": True, "tooltip": "Compile only the transformer blocks, usually enough and can make compilation faster and less error prone"}), + }, + "optional": { + "dynamo_recompile_limit": ("INT", {"default": 128, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.recompile_limit"}), + }, + } + RETURN_TYPES = ("WANCOMPILEARGS",) + RETURN_NAMES = ("torch_compile_args",) + FUNCTION = "set_args" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch 2.5.0 is recommended" + + def set_args(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_transformer_blocks_only, dynamo_recompile_limit=128): + + compile_args = { + "backend": backend, + "fullgraph": fullgraph, + "mode": mode, + "dynamic": dynamic, + "dynamo_cache_size_limit": dynamo_cache_size_limit, + "dynamo_recompile_limit": dynamo_recompile_limit, + "compile_transformer_blocks_only": compile_transformer_blocks_only, + } + + return (compile_args, ) + + class WanVideoLoraSelect: @classmethod def INPUT_TYPES(s): @@ -1052,6 +1137,151 @@ class WanVideoTinyVAELoader: return (vae,) +class LoadWanVideoT5TextEncoder: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": (folder_paths.get_filename_list("text_encoders"), {"tooltip": "These models are loaded from 'ComfyUI/models/text_encoders'"}), + "precision": (["fp32", "bf16"], + {"default": "bf16"} + ), + }, + "optional": { + "load_device": (["main_device", "offload_device"], {"default": "offload_device"}), + "quantization": (['disabled', 'fp8_e4m3fn'], {"default": 'disabled', "tooltip": "optional quantization method"}), + } + } + + RETURN_TYPES = ("WANTEXTENCODER",) + RETURN_NAMES = ("wan_t5_model", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Loads Wan text_encoder model from 'ComfyUI/models/LLM'" + + def loadmodel(self, model_name, precision, load_device="offload_device", quantization="disabled"): + text_encoder_load_device = device if load_device == "main_device" else offload_device + + tokenizer_path = os.path.join(script_directory, "configs", "T5_tokenizer") + + dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + + model_path = folder_paths.get_full_path("text_encoders", model_name) + sd = load_torch_file(model_path, safe_load=True) + + if "token_embedding.weight" not in sd and "shared.weight" not in sd: + raise ValueError("Invalid T5 text encoder model, this node expects the 'umt5-xxl' model") + if "scaled_fp8" in sd: + raise ValueError("Invalid T5 text encoder model, fp8 scaled is not supported by this node") + + # Convert state dict keys from T5 format to the expected format + if "shared.weight" in sd: + log.info("Converting T5 text encoder model to the expected format...") + converted_sd = {} + + for key, value in sd.items(): + # Handle encoder block patterns + if key.startswith('encoder.block.'): + parts = key.split('.') + block_num = parts[2] + + # Self-attention components + if 'layer.0.SelfAttention' in key: + if key.endswith('.k.weight'): + new_key = f"blocks.{block_num}.attn.k.weight" + elif key.endswith('.o.weight'): + new_key = f"blocks.{block_num}.attn.o.weight" + elif key.endswith('.q.weight'): + new_key = f"blocks.{block_num}.attn.q.weight" + elif key.endswith('.v.weight'): + new_key = f"blocks.{block_num}.attn.v.weight" + elif 'relative_attention_bias' in key: + new_key = f"blocks.{block_num}.pos_embedding.embedding.weight" + else: + new_key = key + + # Layer norms + elif 'layer.0.layer_norm' in key: + new_key = f"blocks.{block_num}.norm1.weight" + elif 'layer.1.layer_norm' in key: + new_key = f"blocks.{block_num}.norm2.weight" + + # Feed-forward components + elif 'layer.1.DenseReluDense' in key: + if 'wi_0' in key: + new_key = f"blocks.{block_num}.ffn.gate.0.weight" + elif 'wi_1' in key: + new_key = f"blocks.{block_num}.ffn.fc1.weight" + elif 'wo' in key: + new_key = f"blocks.{block_num}.ffn.fc2.weight" + else: + new_key = key + else: + new_key = key + elif key == "shared.weight": + new_key = "token_embedding.weight" + elif key == "encoder.final_layer_norm.weight": + new_key = "norm.weight" + else: + new_key = key + converted_sd[new_key] = value + sd = converted_sd + + T5_text_encoder = T5EncoderModel( + text_len=512, + dtype=dtype, + device=text_encoder_load_device, + state_dict=sd, + tokenizer_path=tokenizer_path, + quantization=quantization + ) + text_encoder = { + "model": T5_text_encoder, + "dtype": dtype, + } + + return (text_encoder,) + +class LoadWanVideoClipTextEncoder: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": (folder_paths.get_filename_list("clip_vision") + folder_paths.get_filename_list("text_encoders"), {"tooltip": "These models are loaded from 'ComfyUI/models/clip_vision'"}), + "precision": (["fp16", "fp32", "bf16"], + {"default": "fp16"} + ), + }, + "optional": { + "load_device": (["main_device", "offload_device"], {"default": "offload_device"}), + } + } + + RETURN_TYPES = ("CLIP_VISION",) + RETURN_NAMES = ("wan_clip_vision", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Loads Wan clip_vision model from 'ComfyUI/models/clip_vision'" + + def loadmodel(self, model_name, precision, load_device="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] + + model_path = folder_paths.get_full_path("clip_vision", model_name) + # We also support legacy setups where the model is in the text_encoders folder + if model_path is None: + model_path = folder_paths.get_full_path("text_encoders", model_name) + sd = load_torch_file(model_path, safe_load=True) + if "log_scale" not in sd: + raise ValueError("Invalid CLIP model, this node expectes the 'open-clip-xlm-roberta-large-vit-huge-14' model") + + clip_model = CLIPModel(dtype=dtype, device=device, state_dict=sd) + clip_model.model.to(text_encoder_load_device) + del sd + + return (clip_model,) + NODE_CLASS_MAPPINGS = { "WanVideoModelLoader": WanVideoModelLoader, "WanVideoVAELoader": WanVideoVAELoader, @@ -1059,7 +1289,12 @@ NODE_CLASS_MAPPINGS = { "WanVideoLoraBlockEdit": WanVideoLoraBlockEdit, "WanVideoTinyVAELoader": WanVideoTinyVAELoader, "WanVideoVACEModelSelect": WanVideoVACEModelSelect, - "WanVideoLoraSelectMulti": WanVideoLoraSelectMulti + "WanVideoLoraSelectMulti": WanVideoLoraSelectMulti, + "WanVideoBlockSwap": WanVideoBlockSwap, + "WanVideoVRAMManagement": WanVideoVRAMManagement, + "WanVideoTorchCompileSettings": WanVideoTorchCompileSettings, + "LoadWanVideoT5TextEncoder": LoadWanVideoT5TextEncoder, + "LoadWanVideoClipTextEncoder": LoadWanVideoClipTextEncoder, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -1069,5 +1304,10 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoLoraBlockEdit": "WanVideo Lora Block Edit", "WanVideoTinyVAELoader": "WanVideo Tiny VAE Loader", "WanVideoVACEModelSelect": "WanVideo VACE Module Select", - "WanVideoLoraSelectMulti": "WanVideo Lora Select Multi" + "WanVideoLoraSelectMulti": "WanVideo Lora Select Multi", + "WanVideoBlockSwap": "WanVideo Block Swap", + "WanVideoVRAMManagement": "WanVideo VRAM Management", + "WanVideoTorchCompileSettings": "WanVideo Torch Compile Settings", + "LoadWanVideoT5TextEncoder": "WanVideo T5 Text Encoder Loader", + "LoadWanVideoClipTextEncoder": "WanVideo CLIP Text Encoder Loader", } \ No newline at end of file diff --git a/skyreels/nodes.py b/skyreels/nodes.py index d2d8fd8..dceb035 100644 --- a/skyreels/nodes.py +++ b/skyreels/nodes.py @@ -1,7 +1,7 @@ import os import torch import gc -from ..utils import log, print_memory, fourier_filter +from ..utils import log, print_memory, fourier_filter, optimized_scale import math from tqdm import tqdm @@ -9,7 +9,7 @@ from ..wanvideo.modules.model import rope_params from ..wanvideo.schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler from diffusers.schedulers import FlowMatchEulerDiscreteScheduler from ..wanvideo.schedulers.scheduling_flow_match_lcm import FlowMatchLCMScheduler -from ..nodes import optimized_scale + from einops import rearrange from ..enhance_a_video.globals import disable_enhance diff --git a/utils.py b/utils.py index 104b627..daca9d7 100644 --- a/utils.py +++ b/utils.py @@ -241,3 +241,23 @@ def is_image_black(image, threshold=1e-3): if image.min() < 0: image = (image + 1) / 2 return torch.all(image < threshold).item() + +def add_noise_to_reference_video(image, ratio=None): + sigma = torch.ones((image.shape[0],)).to(image.device, image.dtype) * ratio + image_noise = torch.randn_like(image) * sigma[:, None, None, None] + image_noise = torch.where(image==-1, torch.zeros_like(image), image_noise) + image = image + image_noise + return image + +def optimized_scale(positive_flat, negative_flat): + + # Calculate dot production + dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True) + + # Squared norm of uncondition + squared_norm = torch.sum(negative_flat ** 2, dim=1, keepdim=True) + 1e-8 + + # st_star = v_cond^T * v_uncond / ||v_uncond||^2 + st_star = dot_product / squared_norm + + return st_star