krea2: don't hardcode the NVIDIA-only cuDNN SDPA backend (#933)
* krea2: don't hardcode NVIDIA-only cuDNN SDPA backend The krea2 attention() forced SDPBackend.CUDNN_ATTENTION, which is NVIDIA-only. On non-NVIDIA backends (AMD ROCm, Intel XPU, Apple MPS) every forward pass fails with 'RuntimeError: No available kernel. Aborting execution.', so Krea 2 LoRA training cannot run at all there. Pass a priority list [CUDNN, FLASH, EFFICIENT, MATH] instead. NVIDIA still selects cuDNN; other backends fall back to flash/efficient/math. Verified training end-to-end on an AMD Radeon 8060S (gfx1151, ROCm 7.2). * Version bump * Add set priority flag so CUDNN_ATTENTION is selected on cuda devices first. --------- Co-authored-by: Jaret Burkett <jaretburkett@gmail.com>
This commit is contained in:
@@ -56,7 +56,21 @@ def attention(
|
||||
scale: float | None = None,
|
||||
gqa: bool = False,
|
||||
) -> Tensor:
|
||||
with sdpa_kernel(SDPBackend.CUDNN_ATTENTION):
|
||||
# cuDNN attention is NVIDIA-only, so hardcoding SDPBackend.CUDNN_ATTENTION
|
||||
# raises "No available kernel" on non-NVIDIA backends (AMD ROCm, Intel XPU,
|
||||
# Apple MPS). Pass a priority list instead: cuDNN is still preferred on
|
||||
# NVIDIA, and the dispatcher falls back to flash/efficient/math elsewhere.
|
||||
# (On ROCm gfx11xx the flash path needs
|
||||
# TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL=1; masked attention uses math.)
|
||||
with sdpa_kernel(
|
||||
[
|
||||
SDPBackend.CUDNN_ATTENTION,
|
||||
SDPBackend.FLASH_ATTENTION,
|
||||
SDPBackend.EFFICIENT_ATTENTION,
|
||||
SDPBackend.MATH,
|
||||
],
|
||||
set_priority=True,
|
||||
):
|
||||
x = F.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=mask, scale=scale, enable_gqa=gqa
|
||||
)
|
||||
|
||||
@@ -1 +1 @@
|
||||
VERSION = "0.10.24"
|
||||
VERSION = "0.10.25"
|
||||
|
||||
Reference in New Issue
Block a user