Expand k, v when attention backend would fall back to math because gqa. (#15190)

This commit is contained in:
comfyanonymous
2026-07-31 15:17:56 -07:00
committed by GitHub
parent 6cedd34343
commit a1c421994c
2 changed files with 46 additions and 34 deletions

View File

@@ -19,6 +19,7 @@
import torch
import logging
import contextlib
import inspect
import comfy.model_management
from comfy.cli_args import args, PerformanceFeature
import comfy.float
@@ -36,30 +37,59 @@ def run_every_op():
comfy.model_management.throw_exception_if_processing_interrupted()
def gqa_repeat_factor(query_heads, key_heads, value_heads):
if key_heads != value_heads:
raise ValueError(f"Key/value head count mismatch for GQA: {key_heads} != {value_heads}")
if query_heads == key_heads:
return 1
if query_heads % key_heads != 0:
raise ValueError(f"Query heads must be divisible by key/value heads for GQA: {query_heads} vs {key_heads}")
return query_heads // key_heads
def repeat_kv_for_gqa(k, v, query_heads, head_dim):
n_rep = gqa_repeat_factor(query_heads, k.shape[head_dim], v.shape[head_dim])
if n_rep > 1:
k = k.repeat_interleave(n_rep, dim=head_dim)
v = v.repeat_interleave(n_rep, dim=head_dim)
return k, v
def scaled_dot_product_attention(q, k, v, *args, **kwargs):
attn_mask = args[0] if len(args) > 0 else kwargs.get("attn_mask")
if kwargs.get("enable_gqa", False) and attn_mask is not None:
k, v = repeat_kv_for_gqa(k, v, q.shape[-3], -3)
kwargs["enable_gqa"] = False
return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs)
try:
if torch.cuda.is_available():
from torch.nn.attention import SDPBackend, sdpa_kernel
import inspect
if "set_priority" in inspect.signature(sdpa_kernel).parameters:
SDPA_BACKEND_PRIORITY = [
SDPBackend.FLASH_ATTENTION,
SDPBackend.CUDNN_ATTENTION,
SDPBackend.EFFICIENT_ATTENTION,
SDPBackend.MATH,
]
if comfy.model_management.WINDOWS:
SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION)
else:
SDPA_BACKEND_PRIORITY.insert(1, SDPBackend.CUDNN_ATTENTION)
def scaled_dot_product_attention(q, k, v, *args, **kwargs):
if q.nelement() < 1024 * 128: # arbitrary number, for small inputs cudnn attention seems slower
return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs)
attn_mask = args[0] if len(args) > 0 else kwargs.get("attn_mask")
if kwargs.get("enable_gqa", False) and attn_mask is not None and not comfy.model_management.is_nvidia():
k, v = repeat_kv_for_gqa(k, v, q.shape[-3], -3)
kwargs["enable_gqa"] = False
with sdpa_kernel(SDPA_BACKEND_PRIORITY, set_priority=True):
if kwargs.get("enable_gqa", False) and attn_mask is not None and q.shape[-3] != k.shape[-3]:
dropout_p = args[1] if len(args) > 1 else kwargs.get("dropout_p", 0.0)
is_causal = args[2] if len(args) > 2 else kwargs.get("is_causal", False)
params = torch.backends.cuda.SDPAParams(q, k, v, attn_mask, dropout_p, is_causal, True)
supports_native_gqa = (
torch.backends.cuda.can_use_flash_attention(params)
or torch.backends.cuda.can_use_cudnn_attention(params)
or torch.backends.cuda.can_use_efficient_attention(params)
)
if not supports_native_gqa:
k, v = repeat_kv_for_gqa(k, v, q.shape[-3], -3)
kwargs["enable_gqa"] = False
return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs)
else:
logging.warning("Torch version too old to set sdpa backend priority.")