diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 73551dd..c1dd5ed 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -21,12 +21,10 @@ from typing import Any, Callable, Dict, List, Optional, Union, Tuple import torch from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback -from diffusers.configuration_utils import FrozenDict from diffusers.image_processor import VaeImageProcessor from diffusers.schedulers import KarrasDiffusionSchedulers from diffusers.utils import ( - deprecate, logging, replace_example_docstring ) @@ -137,7 +135,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): A scheduler to be used in combination with `unet` to denoise the encoded image latents. """ - # model_cpu_offload_seq = "text_encoder->text_encoder_2->transformer->vae" + model_cpu_offload_seq = "transformer" # _optional_components = ["text_encoder_2"] # _exclude_from_cpu_offload = ["transformer"] # _callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"] @@ -353,6 +351,8 @@ class HunyuanVideoPipeline(DiffusionPipeline): stg_mode: Optional[str] = None, stg_block_idx: Optional[int] = -1, stg_scale: Optional[float] = 0.0, + stg_start_percent: Optional[float] = 0.0, + stg_end_percent: Optional[float] = 1.0, **kwargs, ): r""" @@ -538,16 +538,31 @@ class HunyuanVideoPipeline(DiffusionPipeline): if self.interrupt: continue - # expand the latents if we are doing classifier free guidance - latent_model_input = ( + current_step_percentage = i / len(timesteps) + if self.do_spatio_temporal_guidance: + if stg_start_percent <= current_step_percentage <= stg_end_percent: + stg_enabled = True + if self.do_classifier_free_guidance: + latent_model_input = torch.cat([latents] * 3) + else: + latent_model_input = torch.cat([latents] * 2) + else: + stg_enabled = False + stg_mode = None + stg_block_idx = -1 + prompt_embeds = prompt_embeds[0].unsqueeze(0) + prompt_mask = prompt_mask[0].unsqueeze(0) + prompt_embeds_2 = prompt_embeds_2[0].unsqueeze(0) + latent_model_input = latents + else: + stg_enabled = False + # expand the latents if we are doing classifier free guidance + latent_model_input = ( torch.cat([latents] * 2) - if self.do_classifier_free_guidance and not self.do_spatio_temporal_guidance - else torch.cat([latents] * 3) - if self.do_classifier_free_guidance and self.do_spatio_temporal_guidance - else torch.cat([latents] * 2) - if self.do_spatio_temporal_guidance + if self.do_classifier_free_guidance else latents ) + latent_model_input = self.scheduler.scale_model_input( latent_model_input, t ) @@ -596,7 +611,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): ) + self._stg_scale * ( noise_pred_text - noise_pred_perturb ) - elif self.do_spatio_temporal_guidance: + elif self.do_spatio_temporal_guidance and stg_enabled: noise_pred_text, noise_pred_perturb = noise_pred.chunk(2) noise_pred = noise_pred_text + self._stg_scale * ( noise_pred_text - noise_pred_perturb diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 01a362b..53d29b3 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -717,8 +717,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None # --------------------- Pass through DiT blocks ------------------------ for b, block in enumerate(self.double_blocks): - if b >= 0 and b <= self.double_blocks_to_swap: - mm.soft_empty_cache() + if b <= self.double_blocks_to_swap and self.double_blocks_to_swap > 0: + #mm.soft_empty_cache() block.to(self.main_device) double_block_args = [ img, @@ -733,7 +733,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): ] img, txt = block(*double_block_args) - if b >= 0 and b <= self.double_blocks_to_swap: + if b <= self.double_blocks_to_swap and self.double_blocks_to_swap > 0: #mm.soft_empty_cache() block.to(self.offload_device, non_blocking=True) @@ -741,8 +741,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): x = torch.cat((img, txt), 1) if len(self.single_blocks) > 0: for b, block in enumerate(self.single_blocks): - if b >= 0 and b <= self.single_blocks_to_swap: - mm.soft_empty_cache() + if b <= self.single_blocks_to_swap and self.single_blocks_to_swap > 0: + #mm.soft_empty_cache() block.to(self.main_device) curr_stg_mode = stg_mode if b == stg_block_idx else None single_block_args = [ @@ -759,7 +759,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): ] x = block(*single_block_args) - if b >= 0 and b <= self.single_blocks_to_swap: + if b <= self.single_blocks_to_swap and self.single_blocks_to_swap > 0: #mm.soft_empty_cache() block.to(self.offload_device, non_blocking=True) diff --git a/nodes.py b/nodes.py index dca64c1..18b8606 100644 --- a/nodes.py +++ b/nodes.py @@ -96,6 +96,8 @@ class HyVideoSTG: "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"}), + "stg_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply STG"}), + "stg_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply STG"}), }, } RETURN_TYPES = ("STGARGS",) @@ -116,7 +118,7 @@ class HyVideoModelLoader: "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"}), - "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4"], {"default": 'disabled', "tooltip": "optional quantization method"}), + "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"}), }, "optional": { @@ -155,7 +157,7 @@ class HyVideoModelLoader: base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[base_precision] model_path = folder_paths.get_full_path_or_raise("diffusion_models", model) - sd = load_torch_file(model_path, device=transformer_load_device) + sd = load_torch_file(model_path, device=offload_device) in_channels = out_channels = 16 factor_kwargs = {"device": transformer_load_device, "dtype": base_dtype} @@ -178,38 +180,7 @@ class HyVideoModelLoader: **HUNYUAN_VIDEO_CONFIG, **factor_kwargs ) - - 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]) transformer.eval() - - 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"]) - if "torchao" in quantization: try: from torchao.quantization import ( @@ -217,6 +188,7 @@ class HyVideoModelLoader: fpx_weight_only, float8_dynamic_activation_float8_weight, int8_dynamic_activation_int8_weight, + int8_weight_only, int4_weight_only ) except: @@ -228,22 +200,80 @@ class HyVideoModelLoader: # return isinstance(module, nn.Linear) # return False - if "fp6" in quantization: #slower for some reason on 4090 + if "fp6" in quantization: quant_func = fpx_weight_only(3, 2) elif "int4" in quantization: quant_func = int4_weight_only() - elif "fp8dq" in quantization: #very fast on 4090 when compiled + elif "int8" in quantization: + quant_func = int8_weight_only() + elif "fp8dq" in quantization: quant_func = float8_dynamic_activation_float8_weight() elif 'fp8dqrow' in quantization: from torchao.quantization.quant_api import PerRow quant_func = float8_dynamic_activation_float8_weight(granularity=PerRow()) elif 'int8dq' in quantization: quant_func = int8_dynamic_activation_int8_weight() - - quantize_(transformer, quant_func, device=device) + + log.info(f"Quantizing model with {quant_func}") + + for i, block in enumerate(transformer.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]) + 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"]) + quantize_(block, quant_func) + print(block) + block.to(offload_device) + for i, block in enumerate(transformer.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]) + 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"]) + quantize_(block, quant_func) + for name, param in transformer.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]) manual_offloading = False # to disable manual .to(device) calls log.info(f"Quantized transformer blocks to {quantization}") + for name, param in transformer.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"]) scheduler = FlowMatchDiscreteScheduler( @@ -257,7 +287,7 @@ class HyVideoModelLoader: scheduler=scheduler, progress_bar_config=None ) - + pipeline = { "pipe": pipe, "dtype": base_dtype, @@ -750,6 +780,8 @@ class HyVideoSampler: stg_mode=stg_args["stg_mode"] if stg_args is not None else None, stg_block_idx=stg_args["stg_block_idx"] if stg_args is not None else -1, stg_scale=stg_args["stg_scale"] if stg_args is not None else 0.0, + stg_start_percent=stg_args["stg_start_percent"] if stg_args is not None else 0.0, + stg_end_percent=stg_args["stg_end_percent"] if stg_args is not None else 1.0, ) print_memory(device)