diff --git a/nodes.py b/nodes.py index a1f7043..443aa5a 100644 --- a/nodes.py +++ b/nodes.py @@ -1,9 +1,8 @@ import os import torch import json -from typing import List import gc -from .utils import log, check_diffusers_version, print_memory +from .utils import log, print_memory from diffusers.video_processor import VideoProcessor from .hyvideo.constants import PROMPT_TEMPLATE @@ -20,6 +19,8 @@ from accelerate.utils import set_module_tensor_to_device import folder_paths import comfy.model_management as mm from comfy.utils import load_torch_file +import comfy.model_base +import comfy.latent_formats script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -56,7 +57,7 @@ def get_rotary_pos_embed(transformer, video_length, height, width): if len(rope_sizes) != target_ndim: rope_sizes = [1] * (target_ndim - len(rope_sizes)) + rope_sizes # time axis - + if rope_dim_list is None: rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)] assert ( @@ -71,11 +72,89 @@ def get_rotary_pos_embed(transformer, video_length, height, width): ) return freqs_cos, freqs_sin +def filter_state_dict_by_blocks(state_dict, blocks_mapping): + filtered_dict = {} + + for key in state_dict: + if 'double_blocks.' in key or 'single_blocks.' in key: + block_pattern = key.split('diffusion_model.')[1].split('.', 2)[0:2] + block_key = f'{block_pattern[0]}.{block_pattern[1]}.' + + if block_key in blocks_mapping: + filtered_dict[key] = state_dict[key] + + return filtered_dict + +class HyVideoLoraBlockEdit: + def __init__(self): + self.loaded_lora = None + + @classmethod + def INPUT_TYPES(s): + arg_dict = {} + argument = ("BOOLEAN", {"default": True}) + + for i in range(20): + arg_dict["double_blocks.{}.".format(i)] = argument + + for i in range(40): + arg_dict["single_blocks.{}.".format(i)] = argument + + return {"required": arg_dict} + + RETURN_TYPES = ("SELECTEDBLOCKS", ) + RETURN_NAMES = ("blocks", ) + OUTPUT_TOOLTIPS = ("The modified diffusion model.",) + FUNCTION = "select" + + CATEGORY = "HunyuanVideoWrapper" + + def select(self, **kwargs): + selected_blocks = {k: v for k, v in kwargs.items() if v is True} + print("Selected blocks: ", selected_blocks) + return (selected_blocks,) +class HyVideoLoraSelect: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "lora": (folder_paths.get_filename_list("loras"), + {"tooltip": "LORA models are expected to be in ComfyUI/models/loras with .safetensors extension"}), + "strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.0001, "tooltip": "LORA strength, set to 0.0 to unmerge the LORA"}), + }, + "optional": { + "prev_lora":("HYVIDLORA", {"default": None, "tooltip": "For loading multiple LoRAs"}), + "blocks":("SELECTEDBLOCKS", ), + } + } + + RETURN_TYPES = ("HYVIDLORA",) + RETURN_NAMES = ("lora", ) + FUNCTION = "getlorapath" + CATEGORY = "HunyuanVideoWrapper" + DESCRIPTION = "Select a LoRA model from ComfyUI/models/loras" + + def getlorapath(self, lora, strength, blocks=None, prev_lora=None, fuse_lora=False): + loras_list = [] + + lora = { + "path": folder_paths.get_full_path("loras", lora), + "strength": strength, + "name": lora.split(".")[0], + "fuse_lora": fuse_lora, + "blocks": blocks + } + if prev_lora is not None: + loras_list.extend(prev_lora) + + loras_list.append(lora) + return (loras_list,) + class HyVideoBlockSwap: @classmethod def INPUT_TYPES(s): return { - "required": { + "required": { "double_blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 20, "step": 1, "tooltip": "Number of double blocks to swap"}), "single_blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 40, "step": 1, "tooltip": "Number of single blocks to swap"}), "offload_txt_in": ("BOOLEAN", {"default": False, "tooltip": "Offload txt_in layer"}), @@ -95,7 +174,7 @@ class HyVideoSTG: @classmethod def INPUT_TYPES(s): return { - "required": { + "required": { "stg_mode": (["STG-A", "STG-R"],), "stg_block_idx": ("INT", {"default": 0, "min": -1, "max": 39, "step": 1, "tooltip": "Block index to apply STG"}), "stg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Recommended values are ≤2.0"}), @@ -112,6 +191,33 @@ class HyVideoSTG: def setargs(self, **kwargs): return (kwargs, ) + +class HyVideoModel(comfy.model_base.BaseModel): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.pipeline = {} + + def __getitem__(self, k): + return self.pipeline[k] + + def __setitem__(self, k, v): + self.pipeline[k] = v + + +class HyVideoModelConfig: + def __init__(self, dtype): + self.unet_config = {} + self.unet_extra_config = {} + self.latent_format = comfy.latent_formats.LatentFormat() + self.latent_format.latent_channels = 16 + self.manual_cast_dtype = dtype + self.sampling_settings = {"multiplier": 1.0} + # Don't know what this is. Value taken from ComfyUI Mochi model. + self.memory_usage_factor = 2.0 + # denoiser is handled by extension + self.unet_config["disable_unet_model_creation"] = True + + #region Model loading class HyVideoModelLoader: @classmethod @@ -119,8 +225,8 @@ class HyVideoModelLoader: return { "required": { "model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}), - - "base_precision": (["fp16", "fp32", "bf16"], {"default": "bf16"}), + + "base_precision": (["fp32", "bf16"], {"default": "bf16"}), "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4", "torchao_int8"], {"default": 'disabled', "tooltip": "optional quantization method"}), "load_device": (["main_device", "offload_device"], {"default": "main_device"}), }, @@ -133,6 +239,7 @@ class HyVideoModelLoader: ], {"default": "flash_attn"}), "compile_args": ("COMPILEARGS", ), "block_swap_args": ("BLOCKSWAPARGS", ), + "lora": ("HYVIDLORA", {"default": None}), } } @@ -142,7 +249,7 @@ class HyVideoModelLoader: CATEGORY = "HunyuanVideoWrapper" def loadmodel(self, model, base_precision, load_device, quantization, - compile_args=None, attention_mode="sdpa", block_swap_args=None): + compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None): transformer = None manual_offloading = True if "sage" in attention_mode: @@ -154,7 +261,7 @@ class HyVideoModelLoader: device = mm.get_torch_device() offload_device = mm.unet_offload_device() manual_offloading = True - transformer_load_device = device if load_device == "main_device" else offload_device + transformer_load_device = device if load_device == "main_device" or lora is None else offload_device mm.soft_empty_cache() base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[base_precision] @@ -184,7 +291,74 @@ class HyVideoModelLoader: **factor_kwargs ) transformer.eval() - if "torchao" in quantization: + + comfy_model = HyVideoModel( + HyVideoModelConfig(base_dtype), + model_type=comfy.model_base.ModelType.FLOW, + device=device, + ) + scheduler = FlowMatchDiscreteScheduler( + shift=9.0, + reverse=True, + solver="euler", + ) + pipe = HunyuanVideoPipeline( + transformer=transformer, + scheduler=scheduler, + progress_bar_config=None, + base_dtype=base_dtype + ) + + if not "torchao" in quantization: + log.info("Using accelerate to load and assign model weights to device...") + if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast": + dtype = torch.float8_e4m3fn + else: + dtype = base_dtype + params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"} + for name, param in transformer.named_parameters(): + dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype + set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) + + comfy_model.diffusion_model = transformer + patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) + + if lora is not None: + from comfy.sd import load_lora_for_models + for l in lora: + lora_path = l["path"] + lora_strength = l["strength"] + lora_sd = load_torch_file(lora_path, safe_load=True) + if l["blocks"]: + lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"]) + + #for k in lora_sd.keys(): + # print(k) + + patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0) + + comfy.model_management.load_models_gpu([patcher]) + + if quantization == "fp8_e4m3fn_fast": + from .fp8_optimization import convert_fp8_linear + convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep) + + #compile + if compile_args is not None: + torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] + if compile_args["compile_single_blocks"]: + for i, block in enumerate(patcher.model.diffusion_model.single_blocks): + patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + if compile_args["compile_double_blocks"]: + for i, block in enumerate(patcher.model.diffusion_model.double_blocks): + patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + if compile_args["compile_txt_in"]: + patcher.model.diffusion_model.txt_in = torch.compile(patcher.model.diffusion_model.txt_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + if compile_args["compile_vector_in"]: + patcher.model.diffusion_model.vector_in = torch.compile(patcher.model.diffusion_model.vector_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + if compile_args["compile_final_layer"]: + patcher.model.diffusion_model.final_layer = torch.compile(patcher.model.diffusion_model.final_layer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + elif "torchao" in quantization: try: from torchao.quantization import ( quantize_, @@ -195,14 +369,14 @@ class HyVideoModelLoader: int4_weight_only ) except: - raise ImportError("torchao is not installed, please install torchao to use fp8dq") + raise ImportError("torchao is not installed") # def filter_fn(module: nn.Module, fqn: str) -> bool: # target_submodules = {'attn1', 'ff'} # avoid norm layers, 1.5 at least won't work with quantized norm1 #todo: test other models # if any(sub in fqn for sub in target_submodules): # return isinstance(module, nn.Linear) # return False - + if "fp6" in quantization: quant_func = fpx_weight_only(3, 2) elif "int4" in quantization: @@ -218,93 +392,60 @@ class HyVideoModelLoader: quant_func = int8_dynamic_activation_int8_weight() log.info(f"Quantizing model with {quant_func}") + comfy_model.diffusion_model = transformer + patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) - for i, block in enumerate(transformer.single_blocks): + if lora is not None: + from comfy.sd import load_lora_for_models + for l in lora: + lora_path = l["path"] + lora_strength = l["strength"] + lora_sd = load_torch_file(lora_path, safe_load=True) + patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0) + + comfy.model_management.load_models_gpu([patcher]) + + for i, block in enumerate(patcher.model.diffusion_model.single_blocks): log.info(f"Quantizing single_block {i}") for name, _ in block.named_parameters(prefix=f"single_blocks.{i}"): #print(f"Parameter name: {name}") - set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=base_dtype, value=sd[name]) + set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name]) if compile_args is not None: - transformer.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) quantize_(block, quant_func) print(block) block.to(offload_device) - for i, block in enumerate(transformer.double_blocks): + for i, block in enumerate(patcher.model.diffusion_model.double_blocks): log.info(f"Quantizing double_block {i}") for name, _ in block.named_parameters(prefix=f"double_blocks.{i}"): #print(f"Parameter name: {name}") - set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=base_dtype, value=sd[name]) + set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name]) if compile_args is not None: - transformer.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) quantize_(block, quant_func) - for name, param in transformer.named_parameters(): + for name, param in patcher.model.diffusion_model.named_parameters(): if "single_blocks" not in name and "double_blocks" not in name: - set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=base_dtype, value=sd[name]) - + set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name]) + manual_offloading = False # to disable manual .to(device) calls log.info(f"Quantized transformer blocks to {quantization}") - for name, param in transformer.named_parameters(): + for name, param in patcher.model.diffusion_model.named_parameters(): print(name, param.dtype) #param.data = param.data.to(self.vae_dtype).to(device) - else: - log.info("Using accelerate to load and assign model weights to device...") - if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast": - dtype = torch.float8_e4m3fn - else: - dtype = base_dtype - params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"} - for name, param in transformer.named_parameters(): - dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype - set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) - - - if quantization == "fp8_e4m3fn_fast": - from .fp8_optimization import convert_fp8_linear - convert_fp8_linear(transformer, base_dtype, params_to_keep=params_to_keep) - - #compile - if compile_args is not None: - torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] - if compile_args["compile_single_blocks"]: - for i, block in enumerate(transformer.single_blocks): - transformer.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) - if compile_args["compile_double_blocks"]: - for i, block in enumerate(transformer.double_blocks): - transformer.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) - if compile_args["compile_txt_in"]: - transformer.txt_in = torch.compile(transformer.txt_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) - if compile_args["compile_vector_in"]: - transformer.vector_in = torch.compile(transformer.vector_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) - if compile_args["compile_final_layer"]: - transformer.final_layer = torch.compile(transformer.final_layer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) del sd mm.soft_empty_cache() - scheduler = FlowMatchDiscreteScheduler( - shift=9.0, - reverse=True, - solver="euler", - ) - - pipe = HunyuanVideoPipeline( - transformer=transformer, - scheduler=scheduler, - progress_bar_config=None, - base_dtype=base_dtype - ) - - pipeline = { - "pipe": pipe, - "dtype": base_dtype, - "base_path": model_path, - "model_name": model, - "manual_offloading": manual_offloading, - "quantization": "disabled", - "block_swap_args": block_swap_args, - } - return (pipeline,) - + patcher.model["pipe"] = pipe + patcher.model["dtype"] = base_dtype + patcher.model["base_path"] = model_path + patcher.model["model_name"] = model + patcher.model["manual_offloading"] = manual_offloading + patcher.model["quantization"] = "disabled" + patcher.model["block_swap_args"] = block_swap_args + + return (patcher,) + #region load VAE class HyVideoVAELoader: @@ -329,7 +470,7 @@ class HyVideoVAELoader: DESCRIPTION = "Loads Hunyuan VAE model from 'ComfyUI/models/vae'" def loadmodel(self, model_name, precision, compile_args=None): - + device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -345,21 +486,21 @@ class HyVideoVAELoader: vae.requires_grad_(False) vae.eval() vae.to(device = device, dtype = dtype) - + #compile if compile_args is not None: torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] vae = torch.compile(vae, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) return (vae,) - - + + class HyVideoTorchCompileSettings: @classmethod def INPUT_TYPES(s): return { - "required": { + "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"}), @@ -395,7 +536,7 @@ class HyVideoTorchCompileSettings: } return (compile_args, ) - + #region TextEncode class DownloadAndLoadHyVideoTextEncoder: @@ -405,7 +546,7 @@ class DownloadAndLoadHyVideoTextEncoder: "required": { "llm_model": (["Kijai/llava-llama-3-8b-text-encoder-tokenizer",],), "clip_model": (["disabled","openai/clip-vit-large-patch14",],), - + "precision": (["fp16", "fp32", "bf16"], {"default": "bf16"} ), @@ -424,7 +565,7 @@ class DownloadAndLoadHyVideoTextEncoder: DESCRIPTION = "Loads Hunyuan text_encoder model from 'ComfyUI/models/LLM'" def loadmodel(self, llm_model, clip_model, precision, apply_final_norm=False, hidden_state_skip_layer=2, quantization="disabled"): - + device = mm.get_torch_device() offload_device = mm.unet_offload_device() dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] @@ -473,7 +614,7 @@ class DownloadAndLoadHyVideoTextEncoder: local_dir=base_path, local_dir_use_symlinks=False, ) - + text_encoder = TextEncoder( text_encoder_path=base_path, text_encoder_type="llm", @@ -487,13 +628,13 @@ class DownloadAndLoadHyVideoTextEncoder: dtype=dtype, quantization_config=quantization_config ) - - + + hyvid_text_encoders = { "text_encoder": text_encoder, "text_encoder_2": text_encoder_2, } - + return (hyvid_text_encoders,) class HyVideoCustomPromptTemplate: @@ -516,7 +657,7 @@ class HyVideoCustomPromptTemplate: "crop_start": crop_start, } return (prompt_template_dict,) - + class HyVideoTextEncode: @classmethod def INPUT_TYPES(s): @@ -549,7 +690,7 @@ class HyVideoTextEncode: text_encoder_2 = None negative_prompt = None - + if prompt_template != "disabled": if prompt_template == "custom": prompt_template_dict = custom_prompt_template @@ -574,12 +715,12 @@ class HyVideoTextEncode: batch_size = 1 num_videos_per_prompt = 1 do_classifier_free_guidance = False # not implemented, for now we only have cfg distilled model - + text_inputs = text_encoder.text2tokens(prompt, prompt_template=prompt_template_dict) prompt_outputs = text_encoder.encode(text_inputs, prompt_template=prompt_template_dict, device=device) prompt_embeds = prompt_outputs.hidden_state - + attention_mask = prompt_outputs.attention_mask log.info(f"{text_encoder.text_encoder_type} prompt attention_mask shape: {attention_mask.shape}, masked tokens: {attention_mask[0].sum().item()}") if attention_mask is not None: @@ -696,7 +837,7 @@ class HyVideoTextEncode: prompt_embeds_2 = prompt_embeds_2.to(device=device) negative_prompt_embeds_2, attention_mask_2, negative_attention_mask_2 = None, None, None - + if force_offload: clip_l.cond_stage_model.to(offload_device) mm.soft_empty_cache() @@ -705,7 +846,7 @@ class HyVideoTextEncode: negative_prompt_embeds_2 = None attention_mask_2 = None negative_attention_mask_2 = None - + prompt_embeds_dict = { "prompt_embeds": prompt_embeds, "negative_prompt_embeds": negative_prompt_embeds, @@ -717,9 +858,9 @@ class HyVideoTextEncode: "negative_attention_mask_2": negative_attention_mask_2, } return (prompt_embeds_dict,) - -#region Sampler + +#region Sampler class HyVideoSampler: @classmethod def INPUT_TYPES(s): @@ -735,7 +876,7 @@ class HyVideoSampler: "flow_shift": ("FLOAT", {"default": 9.0, "min": 0.0, "max": 30.0, "step": 0.01}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "force_offload": ("BOOLEAN", {"default": True}), - + }, "optional": { "samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ), @@ -750,6 +891,8 @@ class HyVideoSampler: CATEGORY = "HunyuanVideoWrapper" def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames, samples=None, denoise_strength=1.0, force_offload=True, stg_args=None): + model = model.model + device = mm.get_torch_device() offload_device = mm.unet_offload_device() dtype = model["dtype"] @@ -760,7 +903,7 @@ class HyVideoSampler: raise ValueError( f"STG-A requires attention_mode to be 'sdpa', but got {transformer.attention_mode}." ) - + generator = torch.Generator(device=torch.device("cpu")).manual_seed(seed) if width <= 0 or height <= 0 or num_frames <= 0: @@ -785,7 +928,7 @@ class HyVideoSampler: n_tokens = freqs_cos.shape[0] model["pipe"].scheduler.shift = flow_shift - + # autocast_context = torch.autocast( # mm.get_autocast_device(device), dtype=dtype # ) if any(q in model["quantization"] for q in ("e4m3fn", "GGUF")) else nullcontext() @@ -795,9 +938,9 @@ class HyVideoSampler: #print(name, param.data.device) if "single" not in name and "double" not in name: param.data = param.data.to(device) - + transformer.block_swap( - model["block_swap_args"]["double_blocks_to_swap"] - 1 , + model["block_swap_args"]["double_blocks_to_swap"] - 1 , model["block_swap_args"]["single_blocks_to_swap"] - 1, offload_txt_in = model["block_swap_args"]["offload_txt_in"], offload_img_in = model["block_swap_args"]["offload_img_in"], @@ -808,18 +951,17 @@ class HyVideoSampler: elif model["manual_offloading"]: transformer.to(device) - mm.unload_all_models() mm.soft_empty_cache() gc.collect() - + try: torch.cuda.reset_peak_memory_stats(device) except: pass - + #for name, param in transformer.named_parameters(): # print(name, param.data.device) - + out_latents = model["pipe"]( num_inference_steps=steps, height = target_height, @@ -845,7 +987,7 @@ class HyVideoSampler: torch.cuda.reset_peak_memory_stats(device) except: pass - + if force_offload: if model["manual_offloading"]: transformer.to(offload_device) @@ -855,8 +997,8 @@ class HyVideoSampler: return ({ "samples": out_latents },) - -#region VideoDecode + +#region VideoDecode class HyVideoDecode: @classmethod def INPUT_TYPES(s): @@ -867,7 +1009,7 @@ class HyVideoDecode: "temporal_tiling_sample_size": ("INT", {"default": 64, "min": 4, "max": 256, "tooltip": "Smaller values use less VRAM, model default is 64, any other value will cause stutter"}), "spatial_tile_sample_min_size": ("INT", {"default": 256, "min": 32, "max": 2048, "step": 32, "tooltip": "Spatial tile minimum size in pixels, smaller values use less VRAM, may introduce more seams"}), "auto_tile_size": ("BOOLEAN", {"default": True, "tooltip": "Automatically set tile size based on defaults, above settings are ignored"}), - }, + }, } RETURN_TYPES = ("IMAGE",) @@ -895,8 +1037,8 @@ class HyVideoDecode: vae.tile_latent_min_tsize = 16 vae.tile_sample_min_size = 256 vae.tile_latent_min_size = 32 - - + + expand_temporal_dim = False if len(latents.shape) == 4: if isinstance(vae, AutoencoderKLCausal3D): @@ -940,7 +1082,7 @@ class HyVideoDecode: return (out,) -#region VideoEncode +#region VideoEncode class HyVideoEncode: @classmethod def INPUT_TYPES(s): @@ -951,7 +1093,7 @@ class HyVideoEncode: "temporal_tiling_sample_size": ("INT", {"default": 64, "min": 4, "max": 256, "tooltip": "Smaller values use less VRAM, model default is 64, any other value will cause stutter"}), "spatial_tile_sample_min_size": ("INT", {"default": 256, "min": 32, "max": 2048, "step": 32, "tooltip": "Spatial tile minimum size in pixels, smaller values use less VRAM, may introduce more seams"}), "auto_tile_size": ("BOOLEAN", {"default": True, "tooltip": "Automatically set tile size based on defaults, above settings are ignored"}), - }, + }, } RETURN_TYPES = ("LATENT",) @@ -962,7 +1104,7 @@ class HyVideoEncode: def encode(self, vae, image, enable_vae_tiling, temporal_tiling_sample_size, auto_tile_size, spatial_tile_sample_min_size): device = mm.get_torch_device() offload_device = mm.unet_offload_device() - + generator = torch.Generator(device=torch.device("cpu"))#.manual_seed(seed) vae.to(device) if not auto_tile_size: @@ -985,8 +1127,8 @@ class HyVideoEncode: latents = vae.encode(image).latent_dist.sample(generator) latents = latents * vae.config.scaling_factor vae.to(offload_device) - print("encoded latents shape",latents.shape) - + print("encoded latents shape",latents.shape) + return ({"samples": latents},) @@ -1016,32 +1158,32 @@ class HyVideoLatentPreview: latents = samples["samples"].clone() print("in sample", latents.shape) #latent_rgb_factors =[[-0.02531045419704009, -0.00504800612542497, 0.13293717293982546], [-0.03421835830845858, 0.13996708548892614, -0.07081038680118075], [0.011091819063647063, -0.03372949685846012, -0.0698232210116172], [-0.06276524604742019, -0.09322986677909442, 0.01826383612148913], [0.021290659938126788, -0.07719530444034409, -0.08247812477766273], [0.04401102991215147, -0.0026401932105894754, -0.01410913586718443], [0.08979717602613707, 0.05361221258740831, 0.11501425309699129], [0.04695121980405198, -0.13053491609675175, 0.05025986885867986], [-0.09704684176098193, 0.03397687417738002, -0.1105886644677771], [0.14694697234804935, -0.12316902186157716, 0.04210404546699645], [0.14432470831243552, -0.002580008133591355, -0.08490676947390643], [0.051502750076553944, -0.10071695490292451, -0.01786223610178095], [-0.12503276881774464, 0.08877830923879379, 0.1076584501927316], [-0.020191205513213406, -0.1493425056303128, -0.14289740371758308], [-0.06470138952271293, -0.07410426095060325, 0.00980804676890873], [0.11747671720735695, 0.10916082743849789, -0.12235599365235904]] - latent_rgb_factors = [[-0.41, -0.25, -0.26], - [-0.26, -0.49, -0.24], - [-0.37, -0.54, -0.3], - [-0.04, -0.29, -0.29], - [-0.52, -0.59, -0.39], - [-0.56, -0.6, -0.02], - [-0.53, -0.06, -0.48], - [-0.51, -0.28, -0.18], - [-0.59, -0.1, -0.33], - [-0.56, -0.54, -0.41], - [-0.61, -0.19, -0.5], - [-0.05, -0.25, -0.17], - [-0.23, -0.04, -0.22], - [-0.51, -0.56, -0.43], - [-0.13, -0.4, -0.05], + latent_rgb_factors = [[-0.41, -0.25, -0.26], + [-0.26, -0.49, -0.24], + [-0.37, -0.54, -0.3], + [-0.04, -0.29, -0.29], + [-0.52, -0.59, -0.39], + [-0.56, -0.6, -0.02], + [-0.53, -0.06, -0.48], + [-0.51, -0.28, -0.18], + [-0.59, -0.1, -0.33], + [-0.56, -0.54, -0.41], + [-0.61, -0.19, -0.5], + [-0.05, -0.25, -0.17], + [-0.23, -0.04, -0.22], + [-0.51, -0.56, -0.43], + [-0.13, -0.4, -0.05], [-0.01, -0.01, -0.48]] - + import random random.seed(seed) #latent_rgb_factors = [[random.uniform(min_val, max_val) for _ in range(3)] for _ in range(16)] out_factors = latent_rgb_factors print(latent_rgb_factors) - + #latent_rgb_factors_bias = [0.138, 0.025, -0.299] latent_rgb_factors_bias = [r_bias, g_bias, b_bias] - + latent_rgb_factors = torch.tensor(latent_rgb_factors, device=latents.device, dtype=latents.dtype).transpose(0, 1) latent_rgb_factors_bias = torch.tensor(latent_rgb_factors_bias, device=latents.device, dtype=latents.dtype) @@ -1062,9 +1204,9 @@ class HyVideoLatentPreview: latent_images_min = latent_images.min() latent_images_max = latent_images.max() latent_images = (latent_images - latent_images_min) / (latent_images_max - latent_images_min) - + return (latent_images.float().cpu(), out_factors) - + NODE_CLASS_MAPPINGS = { "HyVideoSampler": HyVideoSampler, "HyVideoDecode": HyVideoDecode, @@ -1078,6 +1220,8 @@ NODE_CLASS_MAPPINGS = { "HyVideoSTG": HyVideoSTG, "HyVideoCustomPromptTemplate": HyVideoCustomPromptTemplate, "HyVideoLatentPreview": HyVideoLatentPreview, + "HyVideoLoraSelect": HyVideoLoraSelect, + "HyVideoLoraBlockEdit": HyVideoLoraBlockEdit, } NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoSampler": "HunyuanVideo Sampler", @@ -1092,4 +1236,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoSTG": "HunyuanVideo STG", "HyVideoCustomPromptTemplate": "HunyuanVideo Custom Prompt Template", "HyVideoLatentPreview": "HunyuanVideo Latent Preview", + "HyVideoLoraSelect": "HunyuanVideo Lora Select", + "HyVideoLoraBlockEdit": "HunyuanVideo Lora Block Edit", }