diff --git a/hyvideo/modules/attention.py b/hyvideo/modules/attention.py index 057d44b..9236f06 100644 --- a/hyvideo/modules/attention.py +++ b/hyvideo/modules/attention.py @@ -1,14 +1,13 @@ -import importlib.metadata import math import torch -import torch.nn as nn import torch.nn.functional as F try: from flash_attn.flash_attn_interface import flash_attn_varlen_func except ImportError: flash_attn_varlen_func = None + try: from sageattention import sageattn_varlen @torch.compiler.disable() @@ -22,11 +21,13 @@ try: max_seqlen_kv, ): return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv) -except: - pass +except ImportError: + sageattn_varlen_func = None + +from comfy.ldm.modules.attention import optimized_attention MEMORY_LAYOUT = { - "flash_attn": ( + "flash_attn_varlen": ( lambda x: x.view(x.shape[0] * x.shape[1], *x.shape[2:]), lambda x: x, ), @@ -38,7 +39,7 @@ MEMORY_LAYOUT = { lambda x: x.transpose(1, 2), lambda x: x.transpose(1, 2), ), - "sageattn": ( + "comfy": ( lambda x: x.transpose(1, 2), lambda x: x.transpose(1, 2), ), @@ -79,6 +80,7 @@ def attention( q, k, v, + heads, mode="flash_attn", drop_rate=0, attn_mask=None, @@ -88,6 +90,7 @@ def attention( max_seqlen_q=None, max_seqlen_kv=None, batch_size=1, + ): """ Perform QKV self attention. @@ -136,6 +139,9 @@ def attention( x = x.view( batch_size, max_seqlen_q, x.shape[-2], x.shape[-1] ) # 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 == "flash_attn": x = flash_attn_varlen_func( q, @@ -182,7 +188,8 @@ def attention( else: raise NotImplementedError(f"Unsupported attention mode: {mode}") - x = post_attn_layout(x) - b, s, a, d = x.shape - out = x.reshape(b, s, -1) - return out + if mode != "comfy": + x = post_attn_layout(x) + b, s, a, d = x.shape + return x.reshape(b, s, -1) + return x diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 69bdde0..6c4c874 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -141,6 +141,7 @@ class MMDoubleStreamBlock(nn.Module): max_seqlen_q: Optional[int] = None, max_seqlen_kv: Optional[int] = None, freqs_cis: tuple = None, + attn_mask: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: ( img_mod1_shift, @@ -189,6 +190,7 @@ class MMDoubleStreamBlock(nn.Module): txt_q, txt_k, txt_v = rearrange( txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num ) + # Apply QK-Norm if needed. txt_q = self.txt_attn_q_norm(txt_q).to(txt_v) txt_k = self.txt_attn_k_norm(txt_k).to(txt_v) @@ -204,12 +206,14 @@ class MMDoubleStreamBlock(nn.Module): q, k, v, + heads = self.heads_num, mode=self.attention_mode, cu_seqlens_q=cu_seqlens_q, cu_seqlens_kv=cu_seqlens_kv, max_seqlen_q=max_seqlen_q, max_seqlen_kv=max_seqlen_kv, batch_size=img_k.shape[0], + attn_mask=attn_mask ) img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :] @@ -322,6 +326,7 @@ class MMSingleStreamBlock(nn.Module): max_seqlen_q: Optional[int] = None, max_seqlen_kv: Optional[int] = None, freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None, + attn_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1) x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale) @@ -355,12 +360,14 @@ class MMSingleStreamBlock(nn.Module): q, k, v, + heads = self.heads_num, mode=self.attention_mode, cu_seqlens_q=cu_seqlens_q, cu_seqlens_kv=cu_seqlens_kv, max_seqlen_q=max_seqlen_q, max_seqlen_kv=max_seqlen_kv, batch_size=x.shape[0], + attn_mask=attn_mask ) # Compute activation in mlp stream, cat again and run second linear layer. @@ -651,6 +658,16 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): max_seqlen_q = img_seq_len + txt_seq_len max_seqlen_kv = max_seqlen_q + # 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 + + # Allow attention to all tokens up to total_len + attn_mask[0, :total_len, :total_len] = True + freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None # --------------------- Pass through DiT blocks ------------------------ for b, block in enumerate(self.double_blocks): @@ -666,6 +683,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): max_seqlen_q, max_seqlen_kv, freqs_cis, + attn_mask ] img, txt = block(*double_block_args) @@ -689,6 +707,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): max_seqlen_q, max_seqlen_kv, (freqs_cos, freqs_sin), + attn_mask ] x = block(*single_block_args) diff --git a/hyvideo/modules/token_refiner.py b/hyvideo/modules/token_refiner.py index 4a8d1df..ed0d323 100644 --- a/hyvideo/modules/token_refiner.py +++ b/hyvideo/modules/token_refiner.py @@ -90,7 +90,7 @@ class IndividualTokenRefinerBlock(nn.Module): k = self.self_attn_k_norm(k).to(v) # Self-Attention - attn = attention(q, k, v, mode="sdpa", attn_mask=attn_mask) + attn = attention(q, k, v, heads = self.heads_num,mode="sdpa", attn_mask=attn_mask) x = x + apply_gate(self.self_attn_proj(attn), gate_msa) diff --git a/nodes.py b/nodes.py index bf24d4e..6a56ed1 100644 --- a/nodes.py +++ b/nodes.py @@ -1,15 +1,11 @@ import os import torch import json -from einops import rearrange -from contextlib import nullcontext from typing import List -from pathlib import Path from .utils import log, check_diffusers_version, print_memory from diffusers.video_processor import VideoProcessor -from .hyvideo.constants import PROMPT_TEMPLATE, NEGATIVE_PROMPT, PRECISION_TO_TYPE -from .hyvideo.vae import load_vae +from .hyvideo.constants import PROMPT_TEMPLATE from .hyvideo.text_encoder import TextEncoder from .hyvideo.utils.data_utils import align_to from .hyvideo.modules.posemb_layers import get_nd_rotary_pos_embed @@ -107,8 +103,9 @@ class HyVideoModelLoader: "optional": { "attention_mode": ([ "sdpa", - "flash_attn", + "flash_attn_varlen", "sageattn_varlen", + "comfy", ], {"default": "flash_attn"}), "compile_args": ("COMPILEARGS", ), "block_swap_args": ("BLOCKSWAPARGS", ), @@ -176,8 +173,6 @@ class HyVideoModelLoader: if quantization == "fp8_e4m3fn_fast": from .fp8_optimization import convert_fp8_linear - if "1.5" in model: - params_to_keep.update({"ff"}) #otherwise NaNs convert_fp8_linear(transformer, base_dtype, params_to_keep=params_to_keep) #compile @@ -466,7 +461,7 @@ class HyVideoTextEncode: def encode_prompt(self, prompt, negative_prompt, text_encoder): batch_size = 1 num_videos_per_prompt = 1 - do_classifier_free_guidance = True + do_classifier_free_guidance = False data_type = "video" text_inputs = text_encoder.text2tokens(prompt, data_type=data_type) @@ -475,6 +470,7 @@ class HyVideoTextEncode: prompt_embeds = prompt_outputs.hidden_state attention_mask = prompt_outputs.attention_mask + print("prompt attention_mask: ", attention_mask.shape) if attention_mask is not None: attention_mask = attention_mask.to(device) bs_embed, seq_len = attention_mask.shape @@ -544,6 +540,9 @@ class HyVideoTextEncode: negative_attention_mask = negative_attention_mask.view( batch_size * num_videos_per_prompt, seq_len ) + else: + negative_prompt_embeds = None + negative_attention_mask = None if do_classifier_free_guidance: # duplicate unconditional embeddings for each generation per prompt, using mps friendly method