From 9fe8a827b274fb3b2699e51783db0d76ec089437 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 5 Jan 2024 12:00:22 -0500 Subject: [PATCH] refactor modeldata --- extensions-builtin/sd-webui-controlnet | 2 +- modules/modeldata.py | 122 +++++++++++++++++++++++++ modules/sd_models.py | 43 +-------- modules/sd_vae.py | 28 ++---- modules/shared.py | 84 ++--------------- 5 files changed, 137 insertions(+), 142 deletions(-) create mode 100644 modules/modeldata.py diff --git a/extensions-builtin/sd-webui-controlnet b/extensions-builtin/sd-webui-controlnet index 16798475b..82e406928 160000 --- a/extensions-builtin/sd-webui-controlnet +++ b/extensions-builtin/sd-webui-controlnet @@ -1 +1 @@ -Subproject commit 16798475b5c203cbbbf0aad234275f0d2ead710b +Subproject commit 82e40692870218b8b1f3842ec11d0b5d7acfdb6e diff --git a/modules/modeldata.py b/modules/modeldata.py new file mode 100644 index 000000000..1a3bae946 --- /dev/null +++ b/modules/modeldata.py @@ -0,0 +1,122 @@ +import sys +import threading +from modules import shared, errors + + +class ModelData: + def __init__(self): + self.sd_model = None + self.sd_refiner = None + self.sd_dict = 'None' + self.initial = True + self.lock = threading.Lock() + + def get_sd_model(self): + from modules.sd_models import reload_model_weights + if self.sd_model is None and shared.opts.sd_model_checkpoint != 'None' and not self.lock.locked(): + with self.lock: + try: + self.sd_model = reload_model_weights(op='model') + self.initial = False + except Exception as e: + shared.log.error("Failed to load stable diffusion model") + errors.display(e, "loading stable diffusion model") + self.sd_model = None + return self.sd_model + + def set_sd_model(self, v): + self.sd_model = v + + def get_sd_refiner(self): + from modules.sd_models import reload_model_weights + if self.sd_refiner is None and shared.opts.sd_model_refiner != 'None' and not self.lock.locked(): + with self.lock: + try: + self.sd_refiner = reload_model_weights(op='refiner') + self.initial = False + except Exception as e: + shared.log.error("Failed to load stable diffusion model") + errors.display(e, "loading stable diffusion model") + self.sd_refiner = None + return self.sd_refiner + + def set_sd_refiner(self, v): + self.sd_refiner = v + + +# provides shared.sd_model field as a property +class Shared(sys.modules[__name__].__class__): + @property + def sd_model(self): + import modules.sd_models # pylint: disable=W0621 + if modules.sd_models.model_data.sd_model is None: + shared.log.debug(f'Model requested: fn={sys._getframe().f_back.f_code.co_name}') # pylint: disable=protected-access + return modules.sd_models.model_data.get_sd_model() + + @sd_model.setter + def sd_model(self, value): + import modules.sd_models # pylint: disable=W0621 + modules.sd_models.model_data.set_sd_model(value) + + @property + def sd_refiner(self): + import modules.sd_models # pylint: disable=W0621 + return modules.sd_models.model_data.get_sd_refiner() + + @sd_refiner.setter + def sd_refiner(self, value): + import modules.sd_models # pylint: disable=W0621 + modules.sd_models.model_data.set_sd_refiner(value) + + @property + def backend(self): + return shared.Backend.ORIGINAL if not shared.cmd_opts.use_openvino and shared.opts.data['sd_backend'] == 'original' else shared.Backend.DIFFUSERS + + @property + def sd_model_type(self): + try: + import modules.sd_models # pylint: disable=W0621 + if modules.sd_models.model_data.sd_model is None: + model_type = 'none' + return model_type + if shared.backend == shared.Backend.ORIGINAL: + model_type = 'ldm' + elif "StableDiffusionXL" in self.sd_model.__class__.__name__: + model_type = 'sdxl' + elif "StableDiffusion" in self.sd_model.__class__.__name__: + model_type = 'sd' + elif "LatentConsistencyModel" in self.sd_model.__class__.__name__: + model_type = 'sd' # lcm is compatible with sd + elif "AnimateDiffPipeline" in self.sd_model.__class__.__name__: + model_type = 'sd' # ad is compatible with sd + elif "Kandinsky" in self.sd_model.__class__.__name__: + model_type = 'kandinsky' + else: + model_type = self.sd_model.__class__.__name__ + except Exception: + model_type = 'unknown' + return model_type + + @property + def sd_refiner_type(self): + try: + import modules.sd_models # pylint: disable=W0621 + if modules.sd_models.model_data.sd_refiner is None: + model_type = 'none' + return model_type + if shared.backend == shared.Backend.ORIGINAL: + model_type = 'ldm' + elif "StableDiffusionXL" in self.sd_refiner.__class__.__name__: + model_type = 'sdxl' + elif "StableDiffusion" in self.sd_refiner.__class__.__name__: + model_type = 'sd' + elif "Kandinsky" in self.sd_refiner.__class__.__name__: + model_type = 'kandinsky' + else: + model_type = self.sd_refiner.__class__.__name__ + except Exception: + model_type = 'unknown' + return model_type + + +model_data = ModelData() diff --git a/modules/sd_models.py b/modules/sd_models.py index 71403f8c3..99ee45b1a 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -5,7 +5,6 @@ import json import time import copy import logging -import threading import contextlib import collections import os.path @@ -25,6 +24,7 @@ from modules import paths, shared, shared_items, shared_state, modelloader, devi from modules.timer import Timer from modules.memstats import memory_stats from modules.paths import models_path, script_path +from modules.modeldata import model_data transformers_logging.set_verbosity_error() @@ -536,47 +536,6 @@ sd1_clip_weight = 'cond_stage_model.transformer.text_model.embeddings.token_embe sd2_clip_weight = 'cond_stage_model.model.transformer.resblocks.0.attn.in_proj_weight' -class ModelData: - def __init__(self): - self.sd_model = None - self.sd_refiner = None - self.sd_dict = 'None' - self.initial = True - self.lock = threading.Lock() - - def get_sd_model(self): - if self.sd_model is None and shared.opts.sd_model_checkpoint != 'None' and not self.lock.locked(): - with self.lock: - try: - self.sd_model = reload_model_weights(op='model') - self.initial = False - except Exception as e: - shared.log.error("Failed to load stable diffusion model") - errors.display(e, "loading stable diffusion model") - self.sd_model = None - return self.sd_model - - def set_sd_model(self, v): - self.sd_model = v - - def get_sd_refiner(self): - if self.sd_refiner is None and shared.opts.sd_model_refiner != 'None' and not self.lock.locked(): - with self.lock: - try: - self.sd_refiner = reload_model_weights(op='refiner') - self.initial = False - except Exception as e: - shared.log.error("Failed to load stable diffusion model") - errors.display(e, "loading stable diffusion model") - self.sd_refiner = None - return self.sd_refiner - - def set_sd_refiner(self, v): - self.sd_refiner = v - -model_data = ModelData() - - def change_backend(): shared.log.info(f'Backend changed: {shared.backend}') shared.log.warning('Full server restart required to apply all changes') diff --git a/modules/sd_vae.py b/modules/sd_vae.py index f3707e3d8..6bca6e3a7 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -1,5 +1,4 @@ import os -import collections import glob from copy import deepcopy import torch @@ -12,7 +11,6 @@ base_vae = None loaded_vae_file = None checkpoint_info = None vae_path = os.path.abspath(os.path.join(paths.models_path, 'VAE')) -checkpoints_loaded = collections.OrderedDict() def get_base_vae(model): @@ -145,31 +143,17 @@ def load_vae_dict(filename): def load_vae(model, vae_file=None, vae_source="unknown-source"): global loaded_vae_file # pylint: disable=global-statement - cache_enabled = shared.opts.sd_vae_checkpoint_cache > 0 if vae_file: try: - if cache_enabled and vae_file in checkpoints_loaded: - # use vae checkpoint cache - shared.log.info(f"Loading VAE: model={get_filename(vae_file)} source={vae_source} cached=True") - store_base_vae(model) - _load_vae_dict(model, checkpoints_loaded[vae_file]) - else: - if not os.path.isfile(vae_file): - shared.log.error(f"VAE not found: model={vae_file} source={vae_source}") - return - store_base_vae(model) - vae_dict_1 = load_vae_dict(vae_file) - _load_vae_dict(model, vae_dict_1) - if cache_enabled: - # cache newly loaded vae - checkpoints_loaded[vae_file] = vae_dict_1.copy() + if not os.path.isfile(vae_file): + shared.log.error(f"VAE not found: model={vae_file} source={vae_source}") + return + store_base_vae(model) + vae_dict_1 = load_vae_dict(vae_file) + _load_vae_dict(model, vae_dict_1) except Exception as e: shared.log.error(f"Loading VAE failed: model={vae_file} source={vae_source} {e}") restore_base_vae(model) - # clean up cache if limit is reached - if cache_enabled: - while len(checkpoints_loaded) > shared.opts.sd_vae_checkpoint_cache + 1: # we need to count the current model - checkpoints_loaded.popitem(last=False) # LRU # If vae used is not in dict, update it # It will be removed on refresh though vae_opt = get_filename(vae_file) diff --git a/modules/shared.py b/modules/shared.py index 135b39766..a40338671 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -294,7 +294,7 @@ options_templates.update(options_section(('sd', "Execution & Models"), { "prompt_mean_norm": OptionInfo(True, "Prompt attention mean normalization"), "comma_padding_backtrack": OptionInfo(20, "Prompt padding for long prompts", gr.Slider, {"minimum": 0, "maximum": 74, "step": 1 }), "sd_checkpoint_cache": OptionInfo(0, "Number of cached models", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1}), - "sd_vae_checkpoint_cache": OptionInfo(0, "Number of cached VAEs", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1}), + "sd_vae_checkpoint_cache": OptionInfo(0, "Number of cached VAEs", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": False}), "sd_disable_ckpt": OptionInfo(False, "Disallow usage of models in ckpt format"), })) @@ -952,82 +952,12 @@ def req(url_addr, headers = None, **kwargs): res = SimpleNamespace(**res) return res -class Shared(sys.modules[__name__].__class__): # this class is here to provide sd_model field as a property, so that it can be created and loaded on demand rather than at program startup. - @property - def sd_model(self): - import modules.sd_models # pylint: disable=W0621 - if modules.sd_models.model_data.sd_model is None: - log.debug(f'Model requested: fn={sys._getframe().f_back.f_code.co_name}') # pylint: disable=protected-access - return modules.sd_models.model_data.get_sd_model() - @sd_model.setter - def sd_model(self, value): - import modules.sd_models # pylint: disable=W0621 - modules.sd_models.model_data.set_sd_model(value) - - @property - def sd_refiner(self): - import modules.sd_models # pylint: disable=W0621 - return modules.sd_models.model_data.get_sd_refiner() - - @sd_refiner.setter - def sd_refiner(self, value): - import modules.sd_models # pylint: disable=W0621 - modules.sd_models.model_data.set_sd_refiner(value) - - @property - def backend(self): - return Backend.ORIGINAL if not cmd_opts.use_openvino and opts.data['sd_backend'] == 'original' else Backend.DIFFUSERS - - @property - def sd_model_type(self): - try: - import modules.sd_models # pylint: disable=W0621 - if modules.sd_models.model_data.sd_model is None: - model_type = 'none' - return model_type - if backend == Backend.ORIGINAL: - model_type = 'ldm' - elif "StableDiffusionXL" in self.sd_model.__class__.__name__: - model_type = 'sdxl' - elif "StableDiffusion" in self.sd_model.__class__.__name__: - model_type = 'sd' - elif "LatentConsistencyModel" in self.sd_model.__class__.__name__: - model_type = 'sd' # lcm is compatible with sd - elif "AnimateDiffPipeline" in self.sd_model.__class__.__name__: - model_type = 'sd' # ad is compatible with sd - elif "Kandinsky" in self.sd_model.__class__.__name__: - model_type = 'kandinsky' - else: - model_type = self.sd_model.__class__.__name__ - except Exception: - model_type = 'unknown' - return model_type - - @property - def sd_refiner_type(self): - try: - import modules.sd_models # pylint: disable=W0621 - if modules.sd_models.model_data.sd_refiner is None: - model_type = 'none' - return model_type - if backend == Backend.ORIGINAL: - model_type = 'ldm' - elif "StableDiffusionXL" in self.sd_refiner.__class__.__name__: - model_type = 'sdxl' - elif "StableDiffusion" in self.sd_refiner.__class__.__name__: - model_type = 'sd' - elif "Kandinsky" in self.sd_refiner.__class__.__name__: - model_type = 'kandinsky' - else: - model_type = self.sd_refiner.__class__.__name__ - except Exception: - model_type = 'unknown' - return model_type - -sd_model = None -sd_refiner = None -sd_model_type = '' -sd_refiner_type = '' +sd_model = None # dummy and overwritten by class +sd_refiner = None # dummy and overwritten by class +sd_model_type = '' # dummy and overwritten by class +sd_refiner_type = '' # dummy and overwritten by class compiled_model_state = None + +from modules.modeldata import Shared # pylint: disable=ungrouped-imports sys.modules[__name__].__class__ = Shared