possibly fix up sdpa

This commit is contained in:
kijai
2024-12-06 03:15:02 +02:00
parent 0654246735
commit 31950d0b55
4 changed files with 45 additions and 20 deletions
+17 -10
View File
@@ -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
+19
View File
@@ -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)
+1 -1
View File
@@ -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)
+8 -9
View File
@@ -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