diff --git a/nodes.py b/nodes.py index 90c0fe5..874f113 100644 --- a/nodes.py +++ b/nodes.py @@ -13,6 +13,7 @@ from .sgm.util import instantiate_from_config from .SUPIR.util import convert_dtype, load_state_dict script_directory = os.path.dirname(os.path.abspath(__file__)) + try: import xformers import xformers.ops @@ -22,14 +23,6 @@ except: XFORMERS_IS_AVAILABLE = False class SUPIR_Upscale: - def __init__(self): - self.current_sdxl_model = None - self.current_supir_model = None - self.current_diffusion_dtype = None - self.current_encoder_dtype = None - self.tiled_vae_state = None - self.tiled_sampling_state = None - upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"] @classmethod @@ -103,8 +96,8 @@ class SUPIR_Upscale: encoder_dtype="auto", batch_size=1): device = mm.get_torch_device() - image = image.to(device) - + mm.unload_all_models() + SUPIR_MODEL_PATH = folder_paths.get_full_path("checkpoints", supir_model) SDXL_MODEL_PATH = folder_paths.get_full_path("checkpoints", sdxl_model) @@ -158,7 +151,7 @@ class SUPIR_Upscale: if not hasattr(self, "model") or self.model is None or self.current_config != custom_config: self.current_config = custom_config self.model = None - mm.unload_all_models() + mm.soft_empty_cache() if use_tiled_sampling: @@ -211,12 +204,12 @@ class SUPIR_Upscale: if use_tiled_vae: self.model.init_tile_vae(encoder_tile_size=encoder_tile_size_pixels, decoder_tile_size=decoder_tile_size_latent) - + image, = ImageScaleBy.upscale(self, image, resize_method, scale_by) B, H, W, C = image.shape new_height = H // 64 * 64 new_width = W // 64 * 64 - image = image.permute(0, 3, 1, 2).contiguous() + image = image.permute(0, 3, 1, 2).contiguous().to(device) resized_image = F.interpolate(image, size=(new_height, new_width), mode='bicubic', align_corners=False) captions_list = [] diff --git a/sgm/util.py b/sgm/util.py index 7627fef..aab8a75 100644 --- a/sgm/util.py +++ b/sgm/util.py @@ -175,14 +175,16 @@ def instantiate_from_config(config): return get_obj_from_str(config["target"])(**config.get("params", dict())) + def get_obj_from_str(string, reload=False, invalidate_cache=True): + package_directory_name = os.path.basename(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) module, cls = string.rsplit(".", 1) if invalidate_cache: importlib.invalidate_caches() if reload: module_imp = importlib.import_module(module) importlib.reload(module_imp) - return getattr(importlib.import_module(module, package='ComfyUI-SUPIR'), cls) + return getattr(importlib.import_module(module, package=package_directory_name), cls) def append_zero(x):