From 18676996d0347c4bf3c3b1828bc66daa40d13e4d Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Tue, 4 Nov 2025 23:31:14 +0000 Subject: [PATCH] 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. --- modules/devices.py | 39 ++++++++++++++++++++++++++++++++++++++- 1 file changed, 38 insertions(+), 1 deletion(-) diff --git a/modules/devices.py b/modules/devices.py index 354543785..5c38222ec 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -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: