Allow using comfy clip_l, fix fp8 fastmode mem use

This commit is contained in:
kijai
2024-12-08 20:34:31 +02:00
parent 11530dcb8e
commit 9a4abbc535
5 changed files with 50 additions and 49 deletions
+6 -14
View File
@@ -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)
@@ -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
+2 -2
View File
@@ -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]
+38 -29
View File
@@ -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"]
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}"
@@ -807,6 +807,15 @@ 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)
+2
View File
@@ -20,3 +20,5 @@ def print_memory(device):
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")
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
#log.info(f"Memory Summary:\n{memory_summary}")