Allow using comfy clip_l, fix fp8 fastmode mem use
This commit is contained in:
+6
-14
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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")
|
||||
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}")
|
||||
Reference in New Issue
Block a user