From e13d384ef75f306fdf540b87e104334f5610a009 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 18 Feb 2025 11:01:04 +0200 Subject: [PATCH] small fixes --- .../diffusion/pipelines/pipeline_hunyuan_video.py | 14 +++++++------- hyvideo/modules/attention.py | 14 ++++++++++++-- hyvideo/modules/models.py | 2 +- nodes.py | 5 +++-- 4 files changed, 23 insertions(+), 12 deletions(-) diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index c24f084..80fe51d 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -727,16 +727,16 @@ class HunyuanVideoPipeline(DiffusionPipeline): ) latent_model_input = torch.cat([latent_model_input, latent_image_input], dim=1) - if embedded_guidance_scale is not None and not cfg_enabled: + if self.do_classifier_free_guidance: guidance_expand = ( - torch.tensor( - [embedded_guidance_scale] * latent_model_input.shape[0], - dtype=self.base_dtype, - device=device, - ) * 1000.0 + torch.tensor([embedded_guidance_scale] * latents.shape[0] * 2, dtype=self.base_dtype, device=device) + * 1000.0 ) else: - guidance_expand = None + guidance_expand = ( + torch.tensor([embedded_guidance_scale] * latents.shape[0], dtype=self.base_dtype, device=device) + * 1000.0 + ) if use_context_schedule: counter = torch.zeros_like(latent_model_input) diff --git a/hyvideo/modules/attention.py b/hyvideo/modules/attention.py index c5930e9..d5c7404 100644 --- a/hyvideo/modules/attention.py +++ b/hyvideo/modules/attention.py @@ -9,7 +9,7 @@ except ImportError: flash_attn_varlen_func = None try: - from sageattention import sageattn_varlen + from sageattention import sageattn_varlen, sageattn @torch.compiler.disable() def sageattn_varlen_func( q, @@ -21,6 +21,9 @@ try: max_seqlen_kv, ): return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv) + @torch.compiler.disable() + def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False): + return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal) except ImportError: sageattn_varlen_func = None @@ -39,6 +42,10 @@ MEMORY_LAYOUT = { lambda x: x.transpose(1, 2), lambda x: x.transpose(1, 2), ), + "sageattn": ( + lambda x: x.transpose(1, 2), + lambda x: x.transpose(1, 2), + ), "comfy": ( lambda x: x.transpose(1, 2), lambda x: x.transpose(1, 2), @@ -174,7 +181,10 @@ def attention( ) # reshape x to [b, s, a, d] elif mode == "comfy": x = optimized_attention(q, k, v, mask=attn_mask, heads=heads, skip_reshape=True) - + elif mode == "sageattn": + if attn_mask is not None and attn_mask.dtype != torch.bool: + attn_mask = attn_mask.to(q.dtype) + x = sageattn_func(q, k, v, attn_mask=attn_mask, dropout_p=drop_rate, is_causal=causal) elif mode == "flash_attn_varlen": x = flash_attn_varlen_func( q, diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 7a4582c..884c898 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -967,7 +967,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): img_seq_len = img.shape[1] max_seqlen_q = max_seqlen_kv = img_seq_len + txt_seq_len - if self.attention_mode == "sdpa" or self.attention_mode == "comfy": + if "varlen" not in self.attention_mode: cu_seqlens_q, cu_seqlens_kv = None, None # Create a square boolean mask filled with False attn_mask = torch.zeros((1, max_seqlen_q, max_seqlen_q), dtype=torch.bool, device=text_mask.device) diff --git a/nodes.py b/nodes.py index faeadfc..fc20db8 100644 --- a/nodes.py +++ b/nodes.py @@ -289,6 +289,7 @@ class HyVideoModelLoader: "sdpa", "flash_attn_varlen", "sageattn_varlen", + "sageattn", "comfy", ], {"default": "flash_attn"}), "compile_args": ("COMPILEARGS", ), @@ -1000,8 +1001,8 @@ class HyVideoCFG: return {"required": { "negative_prompt": ("STRING", {"default": "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion", "multiline": True} ), "cfg": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "guidance scale"} ), - "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ), - "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "End percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ), }, }