diff --git a/backend/attention.py b/backend/attention.py index 8e1b773c..f5965f34 100644 --- a/backend/attention.py +++ b/backend/attention.py @@ -4,7 +4,7 @@ import einops import torch from backend import memory_management -from backend.args import args, SageAttentionFuncs +from backend.args import SageAttentionFuncs, args from modules.errors import display_once if memory_management.xformers_enabled(): @@ -334,6 +334,7 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha if IS_SAGE_2 and args.sage2_function is not SageAttentionFuncs.auto: from functools import partial + import sageattention _function = getattr(sageattention, f"sageattn_qk_int8_pv_{args.sage2_function.value}") diff --git a/backend/nn/qwen.py b/backend/nn/qwen.py index 51eb6636..32c02844 100644 --- a/backend/nn/qwen.py +++ b/backend/nn/qwen.py @@ -1,4 +1,5 @@ # https://github.com/QwenLM/Qwen-Image (Apache 2.0) + import math from typing import Optional, Tuple @@ -7,7 +8,13 @@ import torch.nn as nn import torch.nn.functional as F from einops import repeat -from backend.attention import attention_function +from backend.memory_management import xformers_enabled + +if xformers_enabled(): + from backend.attention import attention_xformers as attention_function +else: + from backend.attention import attention_pytorch as attention_function + from backend.nn.flux import EmbedND from backend.utils import pad_to_patch_size diff --git a/backend/nn/svdq.py b/backend/nn/svdq.py index 5f9c6950..61b5adbb 100644 --- a/backend/nn/svdq.py +++ b/backend/nn/svdq.py @@ -205,8 +205,13 @@ class SVDQT5(torch.nn.Module): # ========== Qwen ========== # +from backend.memory_management import xformers_enabled + +if xformers_enabled(): + from backend.attention import attention_xformers as attention_function +else: + from backend.attention import attention_pytorch as attention_function -from backend.attention import attention_function from backend.nn.flux import EmbedND from backend.nn.qwen import ( GELU,