Allow using comfy text encoding, stop using text_mask to fix normal sageattn and improve sdpa

This commit is contained in:
kijai
2025-03-09 23:33:26 +02:00
parent 796297fc08
commit 7f92dcb5f2
3 changed files with 106 additions and 44 deletions
@@ -541,41 +541,41 @@ class HunyuanVideoPipeline(DiffusionPipeline):
batch_size = 1 batch_size = 1
device = self._execution_device device = self._execution_device
prompt_embeds = prompt_embed_dict["prompt_embeds"] prompt_embeds = prompt_embed_dict.get("prompt_embeds", None)
negative_prompt_embeds = prompt_embed_dict["negative_prompt_embeds"] negative_prompt_embeds = prompt_embed_dict.get("negative_prompt_embeds", None)
prompt_mask = prompt_embed_dict["attention_mask"] #prompt_mask = prompt_embed_dict.get("attention_mask", None)
negative_prompt_mask = prompt_embed_dict["negative_attention_mask"] #negative_prompt_mask = prompt_embed_dict.get("negative_attention_mask", None)
prompt_embeds_2 = prompt_embed_dict["prompt_embeds_2"] prompt_embeds_2 = prompt_embed_dict.get("prompt_embeds_2", None)
negative_prompt_embeds_2 = prompt_embed_dict["negative_prompt_embeds_2"] negative_prompt_embeds_2 = prompt_embed_dict.get("negative_prompt_embeds_2", None)
# For classifier free guidance, we need to do two forward passes. # For classifier free guidance, we need to do two forward passes.
# Here we concatenate the unconditional and text embeddings into a single batch # Here we concatenate the unconditional and text embeddings into a single batch
# to avoid doing two forward passes # to avoid doing two forward passes
if self.do_classifier_free_guidance and not self.do_spatio_temporal_guidance: if self.do_classifier_free_guidance and not self.do_spatio_temporal_guidance:
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds]) prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds])
if prompt_mask is not None: # if prompt_mask is not None:
prompt_mask = torch.cat([negative_prompt_mask, prompt_mask]) # prompt_mask = torch.cat([negative_prompt_mask, prompt_mask])
if prompt_embeds_2 is not None: if prompt_embeds_2 is not None:
prompt_embeds_2 = torch.cat([negative_prompt_embeds_2, prompt_embeds_2]) prompt_embeds_2 = torch.cat([negative_prompt_embeds_2, prompt_embeds_2])
elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance: elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance:
prompt_embeds = torch.cat( prompt_embeds = torch.cat(
[negative_prompt_embeds, prompt_embeds, prompt_embeds] [negative_prompt_embeds, prompt_embeds, prompt_embeds]
) )
if prompt_mask is not None: # if prompt_mask is not None:
prompt_mask = torch.cat([negative_prompt_mask, prompt_mask, prompt_mask]) # prompt_mask = torch.cat([negative_prompt_mask, prompt_mask, prompt_mask])
if prompt_embeds_2 is not None: if prompt_embeds_2 is not None:
prompt_embeds_2 = torch.cat( prompt_embeds_2 = torch.cat(
[negative_prompt_embeds_2, prompt_embeds_2, prompt_embeds_2] [negative_prompt_embeds_2, prompt_embeds_2, prompt_embeds_2]
) )
elif self.do_spatio_temporal_guidance: elif self.do_spatio_temporal_guidance:
prompt_embeds = torch.cat([prompt_embeds, prompt_embeds]) prompt_embeds = torch.cat([prompt_embeds, prompt_embeds])
if prompt_mask is not None: # if prompt_mask is not None:
prompt_mask = torch.cat([prompt_mask, prompt_mask]) # prompt_mask = torch.cat([prompt_mask, prompt_mask])
if prompt_embeds_2 is not None: if prompt_embeds_2 is not None:
prompt_embeds_2 = torch.cat([prompt_embeds_2, prompt_embeds_2]) prompt_embeds_2 = torch.cat([prompt_embeds_2, prompt_embeds_2])
prompt_embeds = prompt_embeds.to(device = device, dtype = self.base_dtype) prompt_embeds = prompt_embeds.to(device = device, dtype = self.base_dtype)
prompt_mask = prompt_mask.to(device) #prompt_mask = prompt_mask.to(device)
if prompt_embeds_2 is not None: if prompt_embeds_2 is not None:
prompt_embeds_2 = prompt_embeds_2.to(device = device, dtype = self.base_dtype) prompt_embeds_2 = prompt_embeds_2.to(device = device, dtype = self.base_dtype)
@@ -695,7 +695,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latent_model_input = latents latent_model_input = latents
input_prompt_embeds = prompt_embeds input_prompt_embeds = prompt_embeds
input_prompt_mask = prompt_mask #input_prompt_mask = prompt_mask
input_prompt_embeds_2 = prompt_embeds_2 input_prompt_embeds_2 = prompt_embeds_2
cfg_enabled = False cfg_enabled = False
stg_enabled = False stg_enabled = False
@@ -712,7 +712,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_mode = None stg_mode = None
stg_block_idx = -1 stg_block_idx = -1
input_prompt_embeds = prompt_embeds[0].unsqueeze(0) input_prompt_embeds = prompt_embeds[0].unsqueeze(0)
input_prompt_mask = prompt_mask[0].unsqueeze(0) #input_prompt_mask = prompt_mask[0].unsqueeze(0)
input_prompt_embeds_2 = prompt_embeds_2[0].unsqueeze(0) input_prompt_embeds_2 = prompt_embeds_2[0].unsqueeze(0)
latent_model_input = latents latent_model_input = latents
else: else:
@@ -726,7 +726,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
cfg_enabled = True cfg_enabled = True
else: else:
input_prompt_embeds = prompt_embeds[1].unsqueeze(0) input_prompt_embeds = prompt_embeds[1].unsqueeze(0)
input_prompt_mask = prompt_mask[1].unsqueeze(0) #input_prompt_mask = prompt_mask[1].unsqueeze(0)
input_prompt_embeds_2 = prompt_embeds_2[1].unsqueeze(0) input_prompt_embeds_2 = prompt_embeds_2[1].unsqueeze(0)
if feta_args is not None: if feta_args is not None:
@@ -799,7 +799,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
partial_latent_model_input, partial_latent_model_input,
t_expand, t_expand,
text_states=input_prompt_embeds, text_states=input_prompt_embeds,
text_mask=input_prompt_mask, #text_mask=input_prompt_mask,
text_states_2=input_prompt_embeds_2, text_states_2=input_prompt_embeds_2,
freqs_cos=freqs_cos, freqs_cos=freqs_cos,
freqs_sin=freqs_sin, freqs_sin=freqs_sin,
@@ -835,7 +835,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latent_model_input, # [2, 16, 33, 24, 42] latent_model_input, # [2, 16, 33, 24, 42]
t_expand, # [2] t_expand, # [2]
text_states=input_prompt_embeds, # [2, 256, 4096] text_states=input_prompt_embeds, # [2, 256, 4096]
text_mask=input_prompt_mask, # [2, 256] #text_mask=input_prompt_mask, # [2, 256]
text_states_2=input_prompt_embeds_2, # [2, 768] text_states_2=input_prompt_embeds_2, # [2, 768]
freqs_cos=freqs_cos, # [seqlen, head_dim] freqs_cos=freqs_cos, # [seqlen, head_dim]
freqs_sin=freqs_sin, # [seqlen, head_dim] freqs_sin=freqs_sin, # [seqlen, head_dim]
@@ -849,7 +849,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latent_model_input[0].unsqueeze(0), latent_model_input[0].unsqueeze(0),
t_expand[0].unsqueeze(0), t_expand[0].unsqueeze(0),
text_states=input_prompt_embeds[0].unsqueeze(0), text_states=input_prompt_embeds[0].unsqueeze(0),
text_mask=input_prompt_mask[0].unsqueeze(0), #text_mask=input_prompt_mask[0].unsqueeze(0),
text_states_2=input_prompt_embeds_2[0].unsqueeze(0), text_states_2=input_prompt_embeds_2[0].unsqueeze(0),
freqs_cos=freqs_cos, freqs_cos=freqs_cos,
freqs_sin=freqs_sin, freqs_sin=freqs_sin,
@@ -862,7 +862,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latent_model_input[1].unsqueeze(0), latent_model_input[1].unsqueeze(0),
t_expand[1].unsqueeze(0), t_expand[1].unsqueeze(0),
text_states=input_prompt_embeds[1].unsqueeze(0), text_states=input_prompt_embeds[1].unsqueeze(0),
text_mask=input_prompt_mask[1].unsqueeze(0), #text_mask=input_prompt_mask[1].unsqueeze(0),
text_states_2=input_prompt_embeds_2[1].unsqueeze(0), text_states_2=input_prompt_embeds_2[1].unsqueeze(0),
freqs_cos=freqs_cos, freqs_cos=freqs_cos,
freqs_sin=freqs_sin, freqs_sin=freqs_sin,
+16 -15
View File
@@ -1043,30 +1043,31 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
if self.offload_img_in: if self.offload_img_in:
self.img_in.to(self.offload_device, non_blocking=True) self.img_in.to(self.offload_device, non_blocking=True)
max_seqlen_q, max_seqlen_kv, attn_mask, cu_seqlens_q, cu_seqlens_kv = None, None, None, None, None
txt_seq_len = txt.shape[1] txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1] img_seq_len = img.shape[1]
max_seqlen_q = max_seqlen_kv = img_seq_len + txt_seq_len
if "varlen" not in self.attention_mode: if "varlen" in self.attention_mode: #just for backwards compatibility
cu_seqlens_q, cu_seqlens_kv = None, None max_seqlen_q = max_seqlen_kv = img_seq_len + txt_seq_len
# Create a square boolean mask filled with False text_mask = torch.ones((1, text_states.shape[1]), dtype=torch.bool, device=text_states.device)
attn_mask = torch.zeros((1, max_seqlen_q, max_seqlen_q), dtype=torch.bool, device=text_mask.device)
# Calculate the valid attention regions
text_len = text_mask[0].sum().item()
total_len = text_len + img_seq_len
# Allow attention to all tokens up to total_len
attn_mask[0, :total_len, :total_len] = True
else:
attn_mask = None
# Compute cu_squlens for flash attention # Compute cu_squlens for flash attention
cu_seqlens_q = get_cu_seqlens(text_mask, img_seq_len) cu_seqlens_q = get_cu_seqlens(text_mask, img_seq_len)
cu_seqlens_kv = cu_seqlens_q cu_seqlens_kv = cu_seqlens_q
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
block_args = [cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv, freqs_cis, attn_mask, self.upcast_rope, token_replace_vec, first_frame_token_num, self.i2v_condition_type] block_args = [
cu_seqlens_q,
cu_seqlens_kv,
max_seqlen_q,
max_seqlen_kv,
freqs_cis,
attn_mask,
self.upcast_rope,
token_replace_vec,
first_frame_token_num,
self.i2v_condition_type
]
#tea_cache #tea_cache
if self.enable_teacache: if self.enable_teacache:
+70 -9
View File
@@ -283,6 +283,7 @@ class HyVideoModelLoader:
"attention_mode": ([ "attention_mode": ([
"sdpa", "sdpa",
"flash_attn_varlen", "flash_attn_varlen",
"sageattn",
"sageattn_varlen", "sageattn_varlen",
"comfy", "comfy",
], {"default": "flash_attn"}), ], {"default": "flash_attn"}),
@@ -302,7 +303,7 @@ class HyVideoModelLoader:
def loadmodel(self, model, base_precision, load_device, quantization, def loadmodel(self, model, base_precision, load_device, quantization,
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, auto_cpu_offload=False, upcast_rope=True): compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, auto_cpu_offload=False, upcast_rope=True):
transformer = None transformer = None
#mm.unload_all_models() mm.unload_all_models()
mm.soft_empty_cache() mm.soft_empty_cache()
manual_offloading = True manual_offloading = True
if "sage" in attention_mode: if "sage" in attention_mode:
@@ -654,6 +655,46 @@ class HyVideoTorchCompileSettings:
return (compile_args, ) return (compile_args, )
#region TextEncode #region TextEncode
class HyVideoTextEmbedBridge:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"positive": ("CONDITIONING", ),
},
"optional": {
"negative": ("CONDITIONING", ),
"hyvid_cfg": ("HYVID_CFG", {"tooltip": "The prompt from the cfg node is not used, only the settings"}),
}
}
RETURN_TYPES = ("HYVIDEMBEDS",)
RETURN_NAMES = ("hyvid_embeds",)
FUNCTION = "convert"
CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "Acts as a bridge between the native ComfyUI conditioning and the HunyuanVideoWrapper embeds"
def convert(self, positive, negative=None, hyvid_cfg=None):
positive_cond = positive[0][0]
positive_pooled = positive[0][1]["pooled_output"]
positive_attention_mask = torch.ones(positive_cond.shape[1], dtype=torch.bool, device=positive_cond.device).unsqueeze(0)
negative_cond, negative_attention_mask, negative_pooled = None, None, None
if negative is not None:
negative_cond = negative[0][0]
negative_pooled = negative[0][1]["pooled_output"]
negative_attention_mask = torch.ones(negative_cond.shape[1], dtype=torch.bool, device=negative_cond.device).unsqueeze(0)
prompt_embeds_dict = {
"prompt_embeds": positive_cond,
"negative_prompt_embeds": negative_cond,
"attention_mask": positive_attention_mask,
"negative_attention_mask": negative_attention_mask,
"prompt_embeds_2": positive_pooled,
"negative_prompt_embeds_2": negative_pooled,
"cfg": torch.tensor(hyvid_cfg["cfg"]) if hyvid_cfg is not None else None,
"start_percent": torch.tensor(hyvid_cfg["start_percent"]) if hyvid_cfg is not None else None,
"end_percent": torch.tensor(hyvid_cfg["end_percent"]) if hyvid_cfg is not None else None,
"batched_cfg": torch.tensor(hyvid_cfg["batched_cfg"]) if hyvid_cfg is not None else None,
}
return (prompt_embeds_dict,)
class DownloadAndLoadHyVideoTextEncoder: class DownloadAndLoadHyVideoTextEncoder:
@classmethod @classmethod
@@ -701,7 +742,7 @@ class DownloadAndLoadHyVideoTextEncoder:
bnb_4bit_quant_type="nf4", bnb_4bit_quant_type="nf4",
bnb_4bit_use_double_quant=True, bnb_4bit_use_double_quant=True,
bnb_4bit_compute_dtype=torch.bfloat16 bnb_4bit_compute_dtype=torch.bfloat16
) )
if clip_model != "disabled": if clip_model != "disabled":
clip_model_path = os.path.join(folder_paths.models_dir, "clip", "clip-vit-large-patch14") clip_model_path = os.path.join(folder_paths.models_dir, "clip", "clip-vit-large-patch14")
@@ -931,10 +972,21 @@ class HyVideoTextEncode:
# max_length = prompt_embeds.shape[1] # max_length = prompt_embeds.shape[1]
uncond_input = text_encoder.text2tokens(uncond_tokens, prompt_template=prompt_template_dict) uncond_input = text_encoder.text2tokens(uncond_tokens, prompt_template=prompt_template_dict)
uncond_image = None
if image is not None:
if text_encoder.text_encoder_type == "vlm":
uncond_image = torch.zeros_like(semantic_images.squeeze(0))
negative_prompt_outputs = text_encoder.encode( negative_prompt_outputs = text_encoder.encode(
uncond_input, prompt_template=prompt_template_dict, device=device uncond_input,
prompt_template=prompt_template_dict,
device=device,
image_token_selection_expr=image_token_selection_expr,
semantic_images = [uncond_image] if text_encoder.text_encoder_type == "vlm" else None,
image_embed_interleave=image_embed_interleave,
data_type=prompt_template,
) )
negative_prompt_embeds = negative_prompt_outputs.hidden_state negative_prompt_embeds = negative_prompt_outputs.hidden_state
negative_attention_mask = negative_prompt_outputs.attention_mask negative_attention_mask = negative_prompt_outputs.attention_mask
@@ -1001,15 +1053,21 @@ class HyVideoTextEncode:
attention_mask_2 = None attention_mask_2 = None
negative_attention_mask_2 = None negative_attention_mask_2 = None
last_token = (attention_mask != 0).sum(dim=1).max().item()
prompt_embeds = prompt_embeds[:, :last_token, :]
if negative_prompt_embeds is not None:
last_token = (negative_attention_mask != 0).sum(dim=1).max().item()
negative_prompt_embeds = negative_prompt_embeds[:, :last_token, :]
prompt_embeds_dict = { prompt_embeds_dict = {
"prompt_embeds": prompt_embeds, "prompt_embeds": prompt_embeds,
"negative_prompt_embeds": negative_prompt_embeds, "negative_prompt_embeds": negative_prompt_embeds,
"attention_mask": attention_mask, #"attention_mask": attention_mask,
"negative_attention_mask": negative_attention_mask, #"negative_attention_mask": negative_attention_mask,
"prompt_embeds_2": prompt_embeds_2, "prompt_embeds_2": prompt_embeds_2,
"negative_prompt_embeds_2": negative_prompt_embeds_2, "negative_prompt_embeds_2": negative_prompt_embeds_2,
"attention_mask_2": attention_mask_2, #"attention_mask_2": attention_mask_2,
"negative_attention_mask_2": negative_attention_mask_2, #"negative_attention_mask_2": negative_attention_mask_2,
"cfg": torch.tensor(hyvid_cfg["cfg"]) if hyvid_cfg is not None else None, "cfg": torch.tensor(hyvid_cfg["cfg"]) if hyvid_cfg is not None else None,
"start_percent": torch.tensor(hyvid_cfg["start_percent"]) if hyvid_cfg is not None else None, "start_percent": torch.tensor(hyvid_cfg["start_percent"]) if hyvid_cfg is not None else None,
"end_percent": torch.tensor(hyvid_cfg["end_percent"]) if hyvid_cfg is not None else None, "end_percent": torch.tensor(hyvid_cfg["end_percent"]) if hyvid_cfg is not None else None,
@@ -1352,6 +1410,7 @@ class HyVideoSampler:
else: else:
transformer.enable_teacache = False transformer.enable_teacache = False
mm.unload_all_models()
mm.soft_empty_cache() mm.soft_empty_cache()
gc.collect() gc.collect()
@@ -1811,7 +1870,8 @@ NODE_CLASS_MAPPINGS = {
"HyVideoTeaCache": HyVideoTeaCache, "HyVideoTeaCache": HyVideoTeaCache,
"HyVideoGetClosestBucketSize": HyVideoGetClosestBucketSize, "HyVideoGetClosestBucketSize": HyVideoGetClosestBucketSize,
"HyVideoI2VEncode": HyVideoI2VEncode, "HyVideoI2VEncode": HyVideoI2VEncode,
"HyVideoEncodeKeyframes": HyVideoEncodeKeyframes "HyVideoEncodeKeyframes": HyVideoEncodeKeyframes,
"HyVideoTextEmbedBridge": HyVideoTextEmbedBridge,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoSampler": "HunyuanVideo Sampler", "HyVideoSampler": "HunyuanVideo Sampler",
@@ -1837,5 +1897,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoTeaCache": "HunyuanVideo TeaCache", "HyVideoTeaCache": "HunyuanVideo TeaCache",
"HyVideoGetClosestBucketSize": "HunyuanVideo Get Closest Bucket Size", "HyVideoGetClosestBucketSize": "HunyuanVideo Get Closest Bucket Size",
"HyVideoI2VEncode": "HyVideo I2V Encode", "HyVideoI2VEncode": "HyVideo I2V Encode",
"HyVideoEncodeKeyframes": "HyVideo Encode Keyframes" "HyVideoEncodeKeyframes": "HyVideo Encode Keyframes",
"HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge",
} }