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
|
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
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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}")
|
||||||
Reference in New Issue
Block a user