From 1dffd114fcd048877f14bc8344510163bb5f9755 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 17 May 2023 15:38:12 -0400 Subject: [PATCH] fix vae loading --- modules/sd_vae.py | 29 +---------------------------- webui.py | 6 +++--- 2 files changed, 4 insertions(+), 31 deletions(-) diff --git a/modules/sd_vae.py b/modules/sd_vae.py index 6b8a9c6f8..7bd294c6d 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -50,7 +50,6 @@ def refresh_vae_list(): global vae_path # pylint: disable=global-statement vae_path = shared.opts.vae_dir vae_dict.clear() - vae_paths = [ os.path.join(sd_models.model_path, '**/*.vae.ckpt'), os.path.join(sd_models.model_path, '**/*.vae.pt'), @@ -74,10 +73,10 @@ def refresh_vae_list(): candidates = [] for path in vae_paths: candidates += glob.iglob(path, recursive=True) - for filepath in candidates: name = get_filename(filepath) vae_dict[name] = filepath + shared.log.info(f"Available VAEs: {vae_path} {len(vae_dict)}") def find_vae_near_checkpoint(checkpoint_file): @@ -92,28 +91,21 @@ def find_vae_near_checkpoint(checkpoint_file): def resolve_vae(checkpoint_file): if shared.cmd_opts.vae is not None: return shared.cmd_opts.vae, 'forced' - is_automatic = shared.opts.sd_vae in {"Automatic", "auto"} # "auto" for people with old config - vae_near_checkpoint = find_vae_near_checkpoint(checkpoint_file) if vae_near_checkpoint is not None and (shared.opts.sd_vae_as_default): return vae_near_checkpoint, 'near checkpoint' - if is_automatic: for named_vae_location in [os.path.join(vae_path, os.path.splitext(os.path.basename(checkpoint_file))[0] + ".vae.pt"), os.path.join(vae_path, os.path.splitext(os.path.basename(checkpoint_file))[0] + ".vae.ckpt"), os.path.join(vae_path, os.path.splitext(os.path.basename(checkpoint_file))[0] + ".vae.safetensors")]: if os.path.isfile(named_vae_location): return named_vae_location, 'in VAE dir' - if shared.opts.sd_vae == "None": return None, None - vae_from_options = vae_dict.get(shared.opts.sd_vae, None) if vae_from_options is not None: return vae_from_options, 'specified in settings' - if not is_automatic: shared.log.warning(f"VAE not found: {shared.opts.sd_vae}") - return None, None @@ -125,10 +117,7 @@ def load_vae_dict(filename): def load_vae(model, vae_file=None, vae_source="from unknown source"): global loaded_vae_file # pylint: disable=global-statement - # save_settings = False - cache_enabled = shared.opts.sd_vae_checkpoint_cache > 0 - if vae_file: if cache_enabled and vae_file in checkpoints_loaded: # use vae checkpoint cache @@ -138,28 +127,22 @@ def load_vae(model, vae_file=None, vae_source="from unknown source"): else: assert os.path.isfile(vae_file), f"VAE {vae_source} doesn't exist: {vae_file}" 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() - # 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) if vae_opt not in vae_dict: vae_dict[vae_opt] = vae_file - elif loaded_vae_file: restore_base_vae(model) - loaded_vae_file = vae_file @@ -179,38 +162,28 @@ unspecified = object() def reload_vae_weights(sd_model=None, vae_file=unspecified): from modules import lowvram, sd_hijack - if not sd_model: sd_model = shared.sd_model - global checkpoint_info # pylint: disable=global-statement checkpoint_info = sd_model.sd_checkpoint_info checkpoint_file = checkpoint_info.filename - if vae_file == unspecified: vae_file, vae_source = resolve_vae(checkpoint_file) else: vae_source = "from function argument" - if loaded_vae_file == vae_file: return - if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: lowvram.send_everything_to_cpu() else: sd_model.to(devices.cpu) - sd_hijack.model_hijack.undo_hijack(sd_model) if shared.cmd_opts.rollback_vae and devices.dtype_vae == torch.bfloat16: devices.dtype_vae = torch.float16 - load_vae(sd_model, vae_file, vae_source) - sd_hijack.model_hijack.hijack(sd_model) script_callbacks.model_loaded_callback(sd_model) - if not shared.cmd_opts.lowvram and not shared.cmd_opts.medvram: sd_model.to(devices.device) - shared.log.info("VAE weights loaded.") return sd_model diff --git a/webui.py b/webui.py index f941375d8..144d33cd2 100644 --- a/webui.py +++ b/webui.py @@ -88,6 +88,9 @@ def initialize(): log.debug('Entering Initialize') check_rollback_vae() + modules.sd_vae.refresh_vae_list() + startup_timer.record("vae") + extensions.list_extensions() startup_timer.record("extensions") @@ -107,9 +110,6 @@ def initialize(): modelloader.load_upscalers() startup_timer.record("upscalers") - modules.sd_vae.refresh_vae_list() - startup_timer.record("vae") - shared.opts.onchange("sd_vae", wrap_queued_call(lambda: modules.sd_vae.reload_vae_weights()), call=False) # shared.opts.onchange("sd_vae_as_default", wrap_queued_call(lambda: modules.sd_vae.reload_vae_weights()), call=False) shared.opts.onchange("temp_dir", ui_tempdir.on_tmpdir_changed)