small fixes
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"} ),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user