From a92f18bf71c46856a7b18d6a2824eec57880ab82 Mon Sep 17 00:00:00 2001 From: DasPauluteli <67437654+DasPauluteli@users.noreply.github.com> Date: Wed, 15 Jul 2026 20:44:34 +0200 Subject: [PATCH] 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 --- .../diffusion_models/krea2/src/mmdit.py | 16 +++++++++++++++- version.py | 2 +- 2 files changed, 16 insertions(+), 2 deletions(-) diff --git a/extensions_built_in/diffusion_models/krea2/src/mmdit.py b/extensions_built_in/diffusion_models/krea2/src/mmdit.py index fec307b..bb6619a 100644 --- a/extensions_built_in/diffusion_models/krea2/src/mmdit.py +++ b/extensions_built_in/diffusion_models/krea2/src/mmdit.py @@ -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 ) diff --git a/version.py b/version.py index c342ff1..7ce0cef 100644 --- a/version.py +++ b/version.py @@ -1 +1 @@ -VERSION = "0.10.24" +VERSION = "0.10.25"