From 5c54bb143702f49f9912b7916f388fc09098467b Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 25 Mar 2024 10:46:36 +0200 Subject: [PATCH] Better xformers check Now relies in the comfy xformers check and can thus be disabled with the command line argument --disable-xformers --- SUPIR/utils/tilevae.py | 18 +++++++++--------- nodes.py | 11 +---------- nodes_v2.py | 16 ++++------------ 3 files changed, 14 insertions(+), 31 deletions(-) diff --git a/SUPIR/utils/tilevae.py b/SUPIR/utils/tilevae.py index de92ef8..c7e2c81 100644 --- a/SUPIR/utils/tilevae.py +++ b/SUPIR/utils/tilevae.py @@ -66,20 +66,20 @@ import torch import torch.version import torch.nn.functional as F from einops import rearrange -#from diffusers.utils.import_utils import is_xformers_available - -#import SUPIR.utils.devices as devices import comfy.model_management device = comfy.model_management.get_torch_device() -try: - import xformers - import xformers.ops - XFORMERS_IS_AVAILABLE = True -except: +if comfy.model_management.XFORMERS_IS_AVAILABLE: + try: + import xformers + import xformers.ops + XFORMERS_IS_AVAILABLE = True + except: + XFORMERS_IS_AVAILABLE = False + print("no module 'xformers'. Processing without...") +else: XFORMERS_IS_AVAILABLE = False - print("no module 'xformers'. Processing without...") sd_flag = True diff --git a/nodes.py b/nodes.py index 2b9a78c..083bb6f 100644 --- a/nodes.py +++ b/nodes.py @@ -21,15 +21,6 @@ from transformers import ( ) script_directory = os.path.dirname(os.path.abspath(__file__)) -try: - import xformers - import xformers.ops - - XFORMERS_IS_AVAILABLE = True -except: - XFORMERS_IS_AVAILABLE = False - - def dummy_build_vision_tower(*args, **kwargs): # Monkey patch the CLIP class before you create an instance. return None @@ -236,7 +227,7 @@ class SUPIR_Upscale: config.model.params.sampler_config.target = f".sgm.modules.diffusionmodules.sampling.{sampler}" print("Using non-tiled sampling") - if XFORMERS_IS_AVAILABLE: + if mm.XFORMERS_IS_AVAILABLE: config.model.params.control_stage_config.params.spatial_transformer_attn_type = "softmax-xformers" config.model.params.network_config.params.spatial_transformer_attn_type = "softmax-xformers" config.model.params.first_stage_config.params.ddconfig.attn_type = "vanilla-xformers" diff --git a/nodes_v2.py b/nodes_v2.py index 94f6af1..23b3595 100644 --- a/nodes_v2.py +++ b/nodes_v2.py @@ -21,15 +21,6 @@ from transformers import ( ) script_directory = os.path.dirname(os.path.abspath(__file__)) -try: - import xformers - import xformers.ops - - XFORMERS_IS_AVAILABLE = True -except: - XFORMERS_IS_AVAILABLE = False - - def dummy_build_vision_tower(*args, **kwargs): # Monkey patch the CLIP class before you create an instance. return None @@ -644,7 +635,8 @@ class SUPIR_model_loader: config = OmegaConf.load(config_path) - if XFORMERS_IS_AVAILABLE: + if mm.XFORMERS_IS_AVAILABLE: + print("Using XFORMERS") config.model.params.control_stage_config.params.spatial_transformer_attn_type = "softmax-xformers" config.model.params.network_config.params.spatial_transformer_attn_type = "softmax-xformers" config.model.params.first_stage_config.params.ddconfig.attn_type = "vanilla-xformers" @@ -784,8 +776,8 @@ class SUPIR_model_loader_v2: mm.soft_empty_cache() config = OmegaConf.load(config_path) - - if XFORMERS_IS_AVAILABLE: + if mm.XFORMERS_IS_AVAILABLE: + print("Using XFORMERS") config.model.params.control_stage_config.params.spatial_transformer_attn_type = "softmax-xformers" config.model.params.network_config.params.spatial_transformer_attn_type = "softmax-xformers" config.model.params.first_stage_config.params.ddconfig.attn_type = "vanilla-xformers"