Spaces:
Running on Zero
Running on Zero
File size: 3,826 Bytes
0ec8119 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 | import torch
import torch.nn.functional as F
from typing import Optional, Tuple
from torch.nn.attention import SDPBackend, sdpa_kernel
from diffusers.models.transformers.transformer_qwenimage import apply_rotary_emb_qwen
class QwenDoubleStreamAttnProcessorFA3:
"""
Attention processor for Qwen double-stream architecture using PyTorch's native
SDPA cuDNN backend (FA3-equivalent fused kernel). Falls back to default SDPA
with a log line if the cuDNN backend fails to dispatch at runtime.
"""
_attention_backend = "cudnn_sdpa"
@torch.no_grad()
def __call__(
self,
attn,
hidden_states: torch.FloatTensor,
encoder_hidden_states: torch.FloatTensor = None,
encoder_hidden_states_mask: torch.FloatTensor = None,
attention_mask: Optional[torch.FloatTensor] = None,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> Tuple[torch.FloatTensor, torch.FloatTensor]:
if encoder_hidden_states is None:
raise ValueError("QwenDoubleStreamAttnProcessorFA3 requires encoder_hidden_states (text stream).")
B, S_img, _ = hidden_states.shape
S_txt = encoder_hidden_states.shape[1]
# QKV projections
img_q = attn.to_q(hidden_states)
img_k = attn.to_k(hidden_states)
img_v = attn.to_v(hidden_states)
txt_q = attn.add_q_proj(encoder_hidden_states)
txt_k = attn.add_k_proj(encoder_hidden_states)
txt_v = attn.add_v_proj(encoder_hidden_states)
# Reshape to (B, S, H, D_h)
H = attn.heads
img_q = img_q.unflatten(-1, (H, -1))
img_k = img_k.unflatten(-1, (H, -1))
img_v = img_v.unflatten(-1, (H, -1))
txt_q = txt_q.unflatten(-1, (H, -1))
txt_k = txt_k.unflatten(-1, (H, -1))
txt_v = txt_v.unflatten(-1, (H, -1))
# Q/K normalization
if getattr(attn, "norm_q", None) is not None:
img_q = attn.norm_q(img_q)
if getattr(attn, "norm_k", None) is not None:
img_k = attn.norm_k(img_k)
if getattr(attn, "norm_added_q", None) is not None:
txt_q = attn.norm_added_q(txt_q)
if getattr(attn, "norm_added_k", None) is not None:
txt_k = attn.norm_added_k(txt_k)
# RoPE
if image_rotary_emb is not None:
img_freqs, txt_freqs = image_rotary_emb
img_q = apply_rotary_emb_qwen(img_q, img_freqs, use_real=False)
img_k = apply_rotary_emb_qwen(img_k, img_freqs, use_real=False)
txt_q = apply_rotary_emb_qwen(txt_q, txt_freqs, use_real=False)
txt_k = apply_rotary_emb_qwen(txt_k, txt_freqs, use_real=False)
# Joint attention: concat along sequence axis -> (B, S_total, H, D_h)
q = torch.cat([txt_q, img_q], dim=1)
k = torch.cat([txt_k, img_k], dim=1)
v = torch.cat([txt_v, img_v], dim=1)
# SDPA expects (B, H, S, D_h)
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
try:
with sdpa_kernel(SDPBackend.CUDNN_ATTENTION):
out = F.scaled_dot_product_attention(q, k, v)
except RuntimeError as e:
print(f"[attn] cuDNN SDPA backend unavailable ({e}), falling back to default", flush=True)
out = F.scaled_dot_product_attention(q, k, v)
# Back to (B, S_total, D_model)
out = out.transpose(1, 2).flatten(2, 3).to(q.dtype)
txt_attn_out = out[:, :S_txt, :]
img_attn_out = out[:, S_txt:, :]
# Output projections
img_attn_out = attn.to_out[0](img_attn_out)
if len(attn.to_out) > 1:
img_attn_out = attn.to_out[1](img_attn_out)
txt_attn_out = attn.to_add_out(txt_attn_out)
return img_attn_out, txt_attn_out
|