fix flash_attn, make sdpa default

This commit is contained in:
kijai
2024-12-06 11:12:10 +02:00
parent d0b98c8dfb
commit 2e8f1f2fbd
6 changed files with 35 additions and 30 deletions
+1 -1
View File
@@ -159,7 +159,7 @@
"bf16",
"fp8_e4m3fn",
"offload_device",
"sageattn_varlen"
"sdpa"
]
},
{
+1 -1
View File
@@ -48,7 +48,7 @@
"bf16",
"fp8_e4m3fn",
"offload_device",
"sageattn_varlen"
"sdpa"
]
},
{
+1 -1
View File
@@ -672,7 +672,7 @@
"bf16",
"fp8_e4m3fn",
"offload_device",
"sageattn_varlen"
"sdpa"
]
}
],
+2 -2
View File
@@ -81,7 +81,7 @@ def attention(
k,
v,
heads,
mode="flash_attn",
mode="sdpa",
drop_rate=0,
attn_mask=None,
causal=False,
@@ -142,7 +142,7 @@ def attention(
elif mode == "comfy":
x = optimized_attention(q, k, v, mask=attn_mask, heads=heads, skip_reshape=True)
elif mode == "flash_attn":
elif mode == "flash_attn_varlen":
x = flash_attn_varlen_func(
q,
k,
+25 -22
View File
@@ -36,7 +36,7 @@ class MMDoubleStreamBlock(nn.Module):
qkv_bias: bool = False,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
attention_mode: str = "flash_attn",
attention_mode: str = "sdpa",
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
@@ -199,9 +199,9 @@ class MMDoubleStreamBlock(nn.Module):
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
v = torch.cat((img_v, txt_v), dim=1)
assert (
cu_seqlens_q.shape[0] == 2 * img.shape[0] + 1
), f"cu_seqlens_q.shape:{cu_seqlens_q.shape}, img.shape[0]:{img.shape[0]}"
#assert (
# cu_seqlens_q.shape[0] == 2 * img.shape[0] + 1
#), f"cu_seqlens_q.shape:{cu_seqlens_q.shape}, img.shape[0]:{img.shape[0]}"
attn = attention(
q,
k,
@@ -262,7 +262,7 @@ class MMSingleStreamBlock(nn.Module):
qk_scale: float = None,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
attention_mode: str = "flash_attn",
attention_mode: str = "sdpa",
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
@@ -353,9 +353,9 @@ class MMSingleStreamBlock(nn.Module):
k = torch.cat((img_k, txt_k), dim=1)
# Compute attention.
assert (
cu_seqlens_q.shape[0] == 2 * x.shape[0] + 1
), f"cu_seqlens_q.shape:{cu_seqlens_q.shape}, x.shape[0]:{x.shape[0]}"
#assert (
# cu_seqlens_q.shape[0] == 2 * x.shape[0] + 1
#), f"cu_seqlens_q.shape:{cu_seqlens_q.shape}, x.shape[0]:{x.shape[0]}"
attn = attention(
q,
k,
@@ -452,7 +452,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
device: Optional[torch.device] = None,
main_device: Optional[torch.device] = None,
offload_device: Optional[torch.device] = None,
attention_mode: str = "flash_attn",
attention_mode: str = "sdpa",
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
@@ -466,6 +466,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
self.main_device = main_device
self.offload_device = offload_device
self.attention_mode = attention_mode
# Text projection. Default to linear projection.
# Alternative: TokenRefiner. See more details (LI-DiT): http://arxiv.org/abs/2406.11831
@@ -651,22 +652,24 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
max_seqlen_q = max_seqlen_kv = img_seq_len + txt_seq_len
# Compute cu_squlens and max_seqlen for flash attention
cu_seqlens_q = get_cu_seqlens(text_mask, img_seq_len)
cu_seqlens_kv = cu_seqlens_q
max_seqlen_q = img_seq_len + txt_seq_len
max_seqlen_kv = max_seqlen_q
if self.attention_mode == "sdpa" or self.attention_mode == "comfy":
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)
# 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)
# Calculate the valid attention regions
text_len = text_mask[0].sum().item()
total_len = text_len + img_seq_len
# 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
# 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
cu_seqlens_q = get_cu_seqlens(text_mask, img_seq_len)
cu_seqlens_kv = cu_seqlens_q
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# --------------------- Pass through DiT blocks ------------------------
+5 -3
View File
@@ -2,6 +2,10 @@
## WORK IN PROGRESS
# Update
Scaled dot product attention (sdpa) should now be working, sageattention is still recommended for speed, but should not be necessary anymore making installation much easier.
Vid2vid test:
[source video](https://www.pexels.com/video/a-4x4-vehicle-speeding-on-a-dirt-road-during-a-competition-15604814/)
@@ -13,8 +17,6 @@ text2vid (old test):
https://github.com/user-attachments/assets/3750da65-9753-4bd2-aae2-a688d2b86115
**Currently seems to require flash_attn (default) or sageattn, spda is not working**
Transformer and VAE (single files, no autodownload):
https://huggingface.co/Kijai/HunyuanVideo_comfy/tree/main
@@ -29,7 +31,7 @@ Files go to `ComfyUI/models/LLM/llava-llama-3-8b-text-encoder-tokenizer`
Clip text encoder (has autodownload)
For now using the original https://huggingface.co/openai/clip-vit-large-patch14, files (only need the .safetensor from the weights) go to:
For now using the original https://huggingface.co/openai/clip-vit-large-patch14, files (only need the .safetensor from the weights and all the config files) go to:
`ComfyUI/models/clip/clip-vit-large-patch14`