FireRed-ImageEdit-Japanese / qwenimage /qwen_fa3_processor.py
someone-in-the-world's picture
Replace vllm-flash-attn3 kernel with native SDPA cuDNN backend; wire up attn processor
0ec8119
Raw
History Blame Contribute Delete
3.83 kB
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