From 9a4abbc5357af643b77f3835b89a2bf4ef93940e Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 8 Dec 2024 20:34:31 +0200 Subject: [PATCH] Allow using comfy clip_l, fix fp8 fastmode mem use --- fp8_optimization.py | 20 ++---- .../pipelines/pipeline_hunyuan_video.py | 2 - hyvideo/modules/models.py | 4 +- nodes.py | 69 +++++++++++-------- utils.py | 4 +- 5 files changed, 50 insertions(+), 49 deletions(-) diff --git a/fp8_optimization.py b/fp8_optimization.py index 09f026d..0688ee6 100644 --- a/fp8_optimization.py +++ b/fp8_optimization.py @@ -7,32 +7,24 @@ def fp8_linear_forward(cls, original_dtype, input): weight_dtype = cls.weight.dtype if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: if len(input.shape) == 3: - if weight_dtype == torch.float8_e4m3fn: - inn = input.reshape(-1, input.shape[2]).to(torch.float8_e5m2) - else: - inn = input.reshape(-1, input.shape[2]).to(torch.float8_e4m3fn) + target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn + inn = input.reshape(-1, input.shape[2]).to(target_dtype) w = cls.weight.t() - scale_weight = torch.ones((1), device=input.device, dtype=torch.float32) - scale_input = scale_weight - + scale = torch.ones((1), device=input.device, dtype=torch.float32) bias = cls.bias.to(original_dtype) if cls.bias is not None else None - out_dtype = original_dtype if bias is not None: - o = torch._scaled_mm(inn, w, out_dtype=out_dtype, bias=bias, scale_a=scale_input, scale_b=scale_weight) + o = torch._scaled_mm(inn, w, out_dtype=original_dtype, bias=bias, scale_a=scale, scale_b=scale) else: - o = torch._scaled_mm(inn, w, out_dtype=out_dtype, scale_a=scale_input, scale_b=scale_weight) + o = torch._scaled_mm(inn, w, out_dtype=original_dtype, scale_a=scale, scale_b=scale) if isinstance(o, tuple): o = o[0] return o.reshape((-1, input.shape[1], cls.weight.shape[0])) else: - cls.to(original_dtype) - out = cls.original_forward(input.to(original_dtype)) - cls.to(original_dtype) - return out + return cls.original_forward(input.to(original_dtype)) else: return cls.original_forward(input) diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index ede5388..51c08d5 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -21,7 +21,6 @@ from typing import Any, Callable, Dict, List, Optional, Union, Tuple import torch from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback -from diffusers.image_processor import VaeImageProcessor from diffusers.schedulers import KarrasDiffusionSchedulers from diffusers.utils import ( @@ -145,7 +144,6 @@ class HunyuanVideoPipeline(DiffusionPipeline): scheduler=scheduler ) self.vae_scale_factor = 8 - self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor) def prepare_extra_func_kwargs(self, func, kwargs): # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 907b180..980f81c 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -706,9 +706,9 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): f"Unsupported text_projection: {self.text_projection}" ) if self.offload_txt_in: - self.txt_in.to(self.offload_device) + self.txt_in.to(self.offload_device, non_blocking=True) if self.offload_img_in: - self.img_in.to(self.offload_device) + self.img_in.to(self.offload_device, non_blocking=True) txt_seq_len = txt.shape[1] img_seq_len = img.shape[1] diff --git a/nodes.py b/nodes.py index c3ab1c9..a633f1e 100644 --- a/nodes.py +++ b/nodes.py @@ -528,6 +528,7 @@ class HyVideoTextEncode: "force_offload": ("BOOLEAN", {"default": True}), "prompt_template": (["video", "image", "custom", "disabled"], {"default": "video", "tooltip": "Use the default prompt templates for the llm text encoder"}), "custom_prompt_template": ("PROMPT_TEMPLATE", {"default": PROMPT_TEMPLATE["dit-llm-encode-video"], "multiline": True}), + "clip_l": ("CLIP", {"tooltip": "Use comfy clip model instead, in this case the text encoder loader's clip_l should be disabled"}), } } @@ -536,12 +537,15 @@ class HyVideoTextEncode: FUNCTION = "process" CATEGORY = "HunyuanVideoWrapper" - def process(self, text_encoders, prompt, force_offload=True, prompt_template="video", custom_prompt_template=None): + def process(self, text_encoders, prompt, force_offload=True, prompt_template="video", custom_prompt_template=None, clip_l=None): device = mm.text_encoder_device() offload_device = mm.text_encoder_offload_device() text_encoder_1 = text_encoders["text_encoder"] - text_encoder_2 = text_encoders["text_encoder_2"] + if clip_l is None: + text_encoder_2 = text_encoders["text_encoder_2"] + else: + text_encoder_2 = None negative_prompt = None @@ -585,27 +589,20 @@ class HyVideoTextEncode: bs_embed * num_videos_per_prompt, seq_len ) - if text_encoder is not None: - prompt_embeds_dtype = text_encoder.dtype - elif self.transformer is not None: - prompt_embeds_dtype = self.transformer.dtype - else: - prompt_embeds_dtype = prompt_embeds.dtype + prompt_embeds = prompt_embeds.to(dtype=text_encoder.dtype, device=device) - prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device) - - if prompt_embeds.ndim == 2: - bs_embed, _ = prompt_embeds.shape - # duplicate text embeddings for each generation per prompt, using mps friendly method - prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt) - prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1) - else: - bs_embed, seq_len, _ = prompt_embeds.shape - # duplicate text embeddings for each generation per prompt, using mps friendly method - prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) - prompt_embeds = prompt_embeds.view( - bs_embed * num_videos_per_prompt, seq_len, -1 - ) + # if prompt_embeds.ndim == 2: + # bs_embed, _ = prompt_embeds.shape + # # duplicate text embeddings for each generation per prompt, using mps friendly method + # prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt) + # prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1) + # else: + # bs_embed, seq_len, _ = prompt_embeds.shape + # # duplicate text embeddings for each generation per prompt, using mps friendly method + # prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) + # prompt_embeds = prompt_embeds.view( + # bs_embed * num_videos_per_prompt, seq_len, -1 + # ) # get unconditional embeddings for classifier free guidance # if do_classifier_free_guidance: @@ -691,6 +688,17 @@ class HyVideoTextEncode: if force_offload: text_encoder_2.to(offload_device) mm.soft_empty_cache() + elif clip_l is not None: + clip_l.cond_stage_model.to(device) + tokens = clip_l.tokenize(prompt, return_word_ids=True) + prompt_embeds_2 = clip_l.encode_from_tokens(tokens, return_pooled=True, return_dict=False)[1] + 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() else: prompt_embeds_2 = None negative_prompt_embeds_2 = None @@ -741,9 +749,6 @@ 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): - mm.unload_all_models() - mm.soft_empty_cache() - device = mm.get_torch_device() offload_device = mm.unet_offload_device() dtype = model["dtype"] @@ -757,11 +762,6 @@ class HyVideoSampler: generator = torch.Generator(device=torch.device("cpu")).manual_seed(seed) - try: - torch.cuda.reset_peak_memory_stats(device) - except: - pass - if width <= 0 or height <= 0 or num_frames <= 0: raise ValueError( f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={num_frames}" @@ -806,7 +806,16 @@ class HyVideoSampler: gc.collect() 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) diff --git a/utils.py b/utils.py index 7d3412c..ac263e9 100644 --- a/utils.py +++ b/utils.py @@ -19,4 +19,6 @@ def print_memory(device): max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3 log.info(f"Allocated memory: {memory=:.3f} GB") log.info(f"Max allocated memory: {max_memory=:.3f} GB") - log.info(f"Max reserved memory: {max_reserved=:.3f} GB") \ No newline at end of file + log.info(f"Max reserved memory: {max_reserved=:.3f} GB") + #memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False) + #log.info(f"Memory Summary:\n{memory_summary}") \ No newline at end of file