mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
Sage Attention 2 + Triton workaround Qwen-Image
Workaround to prevent black images generated with Qwen-Image models when Sage Attention 2 is enabled with Triton as backend on devices with compute capability 8.0 and 8.6. Simply switches back to Cuda backend for these models only. Proof of concept, feel free to close if this is not appropriate.
This commit is contained in:
committed by
GitHub
parent
65dfc9b4d0
commit
18676996d0
+38
-1
@@ -557,14 +557,51 @@ def set_sdpa_params():
|
||||
if 'Sage attention' in opts.sdp_options:
|
||||
try:
|
||||
install('sageattention')
|
||||
from sageattention import sageattn
|
||||
from sageattention import sageattn, sageattn_qk_int8_pv_fp16_cuda
|
||||
|
||||
# Detect GPU architecture - sm80/sm86 need CUDA backend workaround
|
||||
# See: https://github.com/comfyanonymous/ComfyUI/issues/9773
|
||||
use_cuda_backend = False
|
||||
capability = None
|
||||
if torch.cuda.is_available():
|
||||
capability = torch.cuda.get_device_capability()
|
||||
# sm80 = compute capability 8.0 (A100/A6000), sm86 = 8.6 (RTX 3090/3090 Ti)
|
||||
if capability in [(8, 0), (8, 6)]:
|
||||
use_cuda_backend = True
|
||||
log.debug(f'Sage Attention: sm{capability[0]}{capability[1]} GPU detected, will use CUDA backend for Qwen models')
|
||||
|
||||
backend_logged_model = None # Track which model type we've logged for
|
||||
|
||||
sdpa_pre_sage_atten = torch.nn.functional.scaled_dot_product_attention
|
||||
@wraps(sdpa_pre_sage_atten)
|
||||
def sdpa_sage_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.FloatTensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor:
|
||||
nonlocal backend_logged_model
|
||||
if (query.shape[-1] in {128, 96, 64}) and (attn_mask is None) and (query.dtype != torch.float32):
|
||||
if enable_gqa:
|
||||
key = key.repeat_interleave(query.size(-3)//key.size(-3), -3)
|
||||
value = value.repeat_interleave(query.size(-3)//value.size(-3), -3)
|
||||
|
||||
# Use CUDA backend for sm80/sm86 + Qwen models to avoid Triton backend bugs
|
||||
# Other models on sm80/sm86 can use the better Triton backend
|
||||
if use_cuda_backend:
|
||||
from modules import shared
|
||||
if shared.sd_model_type == 'qwen':
|
||||
if backend_logged_model != 'qwen':
|
||||
log.debug(f'Sage Attention: using CUDA backend for Qwen model on sm{capability[0]}{capability[1]} GPU (Triton backend workaround)')
|
||||
backend_logged_model = 'qwen'
|
||||
return sageattn_qk_int8_pv_fp16_cuda(
|
||||
q=query, k=key, v=value,
|
||||
tensor_layout="HND",
|
||||
is_causal=is_causal,
|
||||
sm_scale=scale,
|
||||
return_lse=False,
|
||||
pv_accum_dtype="fp32"
|
||||
)
|
||||
else:
|
||||
if backend_logged_model != shared.sd_model_type:
|
||||
log.debug(f'Sage Attention: using Triton backend (auto-dispatch) for {shared.sd_model_type} model on sm{capability[0]}{capability[1]} GPU')
|
||||
backend_logged_model = shared.sd_model_type
|
||||
# Use normal sageattn (auto-dispatch) for all other cases
|
||||
return sageattn(q=query, k=key, v=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale)
|
||||
else:
|
||||
if enable_gqa:
|
||||
|
||||
Reference in New Issue
Block a user