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