possibly fix up sdpa
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user