small fixes

This commit is contained in:
kijai
2025-02-18 11:01:04 +02:00
parent e063693a0c
commit e13d384ef7
4 changed files with 23 additions and 12 deletions
@@ -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)
+12 -2
View File
@@ -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,
+1 -1
View File
@@ -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)
+3 -2
View File
@@ -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"} ),
},
}