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 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
+19
View File
@@ -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)
+1 -1
View File
@@ -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)
+8 -9
View File
@@ -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