diff --git a/examples/hyvideo_lowvram_blockswap_test.json b/examples/hyvideo_lowvram_blockswap_test.json index 8968198..a74524c 100644 --- a/examples/hyvideo_lowvram_blockswap_test.json +++ b/examples/hyvideo_lowvram_blockswap_test.json @@ -159,7 +159,7 @@ "bf16", "fp8_e4m3fn", "offload_device", - "sageattn_varlen" + "sdpa" ] }, { diff --git a/examples/hyvideo_t2v_example_01.json b/examples/hyvideo_t2v_example_01.json index 4d8df8f..5e13894 100644 --- a/examples/hyvideo_t2v_example_01.json +++ b/examples/hyvideo_t2v_example_01.json @@ -48,7 +48,7 @@ "bf16", "fp8_e4m3fn", "offload_device", - "sageattn_varlen" + "sdpa" ] }, { diff --git a/examples/hyvideo_v2v_example_01.json b/examples/hyvideo_v2v_example_01.json index a4eaf82..5524a5d 100644 --- a/examples/hyvideo_v2v_example_01.json +++ b/examples/hyvideo_v2v_example_01.json @@ -672,7 +672,7 @@ "bf16", "fp8_e4m3fn", "offload_device", - "sageattn_varlen" + "sdpa" ] } ], diff --git a/hyvideo/modules/attention.py b/hyvideo/modules/attention.py index 9236f06..c69fe3c 100644 --- a/hyvideo/modules/attention.py +++ b/hyvideo/modules/attention.py @@ -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, diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 6c4c874..6402979 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -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 ------------------------ diff --git a/readme.md b/readme.md index 6b33e4f..72fc7da 100644 --- a/readme.md +++ b/readme.md @@ -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`