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 weight_dtype = cls.weight.dtype
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
if len(input.shape) == 3: if len(input.shape) == 3:
if weight_dtype == 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(torch.float8_e5m2) inn = input.reshape(-1, input.shape[2]).to(target_dtype)
else:
inn = input.reshape(-1, input.shape[2]).to(torch.float8_e4m3fn)
w = cls.weight.t() w = cls.weight.t()
scale_weight = torch.ones((1), device=input.device, dtype=torch.float32) scale = torch.ones((1), device=input.device, dtype=torch.float32)
scale_input = scale_weight
bias = cls.bias.to(original_dtype) if cls.bias is not None else None bias = cls.bias.to(original_dtype) if cls.bias is not None else None
out_dtype = original_dtype
if bias is not None: 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: 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): if isinstance(o, tuple):
o = o[0] o = o[0]
return o.reshape((-1, input.shape[1], cls.weight.shape[0])) return o.reshape((-1, input.shape[1], cls.weight.shape[0]))
else: else:
cls.to(original_dtype) return cls.original_forward(input.to(original_dtype))
out = cls.original_forward(input.to(original_dtype))
cls.to(original_dtype)
return out
else: else:
return cls.original_forward(input) return cls.original_forward(input)
@@ -21,7 +21,6 @@ from typing import Any, Callable, Dict, List, Optional, Union, Tuple
import torch import torch
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.image_processor import VaeImageProcessor
from diffusers.schedulers import KarrasDiffusionSchedulers from diffusers.schedulers import KarrasDiffusionSchedulers
from diffusers.utils import ( from diffusers.utils import (
@@ -145,7 +144,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
scheduler=scheduler scheduler=scheduler
) )
self.vae_scale_factor = 8 self.vae_scale_factor = 8
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
def prepare_extra_func_kwargs(self, func, kwargs): def prepare_extra_func_kwargs(self, func, kwargs):
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature # 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}" f"Unsupported text_projection: {self.text_projection}"
) )
if self.offload_txt_in: 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: 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] txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1] img_seq_len = img.shape[1]
+39 -30
View File
@@ -528,6 +528,7 @@ class HyVideoTextEncode:
"force_offload": ("BOOLEAN", {"default": True}), "force_offload": ("BOOLEAN", {"default": True}),
"prompt_template": (["video", "image", "custom", "disabled"], {"default": "video", "tooltip": "Use the default prompt templates for the llm text encoder"}), "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}), "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" FUNCTION = "process"
CATEGORY = "HunyuanVideoWrapper" 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() device = mm.text_encoder_device()
offload_device = mm.text_encoder_offload_device() offload_device = mm.text_encoder_offload_device()
text_encoder_1 = text_encoders["text_encoder"] 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 negative_prompt = None
@@ -585,27 +589,20 @@ class HyVideoTextEncode:
bs_embed * num_videos_per_prompt, seq_len bs_embed * num_videos_per_prompt, seq_len
) )
if text_encoder is not None: prompt_embeds = prompt_embeds.to(dtype=text_encoder.dtype, device=device)
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=prompt_embeds_dtype, device=device) # if prompt_embeds.ndim == 2:
# bs_embed, _ = prompt_embeds.shape
if prompt_embeds.ndim == 2: # # duplicate text embeddings for each generation per prompt, using mps friendly method
bs_embed, _ = prompt_embeds.shape # prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt)
# duplicate text embeddings for each generation per prompt, using mps friendly method # prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1)
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt) # else:
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1) # bs_embed, seq_len, _ = prompt_embeds.shape
else: # # duplicate text embeddings for each generation per prompt, using mps friendly method
bs_embed, seq_len, _ = prompt_embeds.shape # prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
# duplicate text embeddings for each generation per prompt, using mps friendly method # prompt_embeds = prompt_embeds.view(
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) # bs_embed * num_videos_per_prompt, seq_len, -1
prompt_embeds = prompt_embeds.view( # )
bs_embed * num_videos_per_prompt, seq_len, -1
)
# get unconditional embeddings for classifier free guidance # get unconditional embeddings for classifier free guidance
# if do_classifier_free_guidance: # if do_classifier_free_guidance:
@@ -691,6 +688,17 @@ class HyVideoTextEncode:
if force_offload: if force_offload:
text_encoder_2.to(offload_device) text_encoder_2.to(offload_device)
mm.soft_empty_cache() 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: else:
prompt_embeds_2 = None prompt_embeds_2 = None
negative_prompt_embeds_2 = None negative_prompt_embeds_2 = None
@@ -741,9 +749,6 @@ class HyVideoSampler:
CATEGORY = "HunyuanVideoWrapper" 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): 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() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
dtype = model["dtype"] dtype = model["dtype"]
@@ -757,11 +762,6 @@ class HyVideoSampler:
generator = torch.Generator(device=torch.device("cpu")).manual_seed(seed) 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: if width <= 0 or height <= 0 or num_frames <= 0:
raise ValueError( raise ValueError(
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={num_frames}" 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() gc.collect()
elif model["manual_offloading"]: elif model["manual_offloading"]:
transformer.to(device) 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(): #for name, param in transformer.named_parameters():
# print(name, param.data.device) # print(name, param.data.device)
+3 -1
View File
@@ -19,4 +19,6 @@ def print_memory(device):
max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3 max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3
log.info(f"Allocated memory: {memory=:.3f} GB") log.info(f"Allocated memory: {memory=:.3f} GB")
log.info(f"Max allocated memory: {max_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}")