diff --git a/modules/processing.py b/modules/processing.py index 2bf4b1f9..399f7e1b 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -819,7 +819,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed: if sd_models.checkpoint_aliases.get(p.override_settings.get("sd_model_checkpoint")) is None: p.override_settings.pop("sd_model_checkpoint", None) - _vae_override = p.override_settings.pop("sd_vae", None) + _vae_override: tuple[str, list[str]] = p.override_settings.pop("sd_vae", None) # apply any options overrides set_config(p.override_settings, is_api=True, run_callbacks=False, save_config=False) @@ -830,7 +830,19 @@ def process_images(p: StableDiffusionProcessing) -> Processed: pass else: manage_model_and_prompt_cache(p) - sd_vae.reload_vae_weights(_vae_override) + if _vae_override is not None: + override, choices = _vae_override + _orig: list[str] = shared.opts.forge_additional_modules.copy() + for i in range(len(_orig)): + if os.path.basename(_orig[i]) in choices: + if _orig[i] != override: + shared.opts.forge_additional_modules.pop(i) + else: + override = None + break + + if sd_vae.reload_vae_weights(override): + shared.opts.forge_additional_modules.append(override) # backwards compatibility, fix sampler and scheduler if invalid sd_samplers.fix_p_invalid_sampler_and_scheduler(p) @@ -844,6 +856,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed: set_config(stored_opts, save_config=False) if _vae_override is not None: sd_vae.restore_vae_weights() + shared.opts.forge_additional_modules = _orig return res diff --git a/modules/sd_vae.py b/modules/sd_vae.py index 504c33d0..9827d5d8 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -36,12 +36,6 @@ def get_loaded_vae_hash() -> str: return sha256[0:10] if sha256 else None -# def get_base_vae(model): -# if base_vae is not None and checkpoint_info == model.sd_checkpoint_info and model: -# return base_vae -# return None - - def store_base_vae(model): global base_vae, checkpoint_info assert loaded_vae_file is None @@ -96,92 +90,14 @@ def refresh_vae_list(): vae_dict.update(dict(sorted(vae_dict.items(), key=lambda item: shared.natural_sort_key(item[0])))) -# def find_vae_near_checkpoint(checkpoint_file): -# checkpoint_path = os.path.basename(checkpoint_file).rsplit(".", 1)[0] -# for vae_file in vae_dict.values(): -# if os.path.basename(vae_file).startswith(checkpoint_path): -# return vae_file - -# return None - - -# @dataclass -# class VaeResolution: -# vae: str = None -# source: str = None -# resolved: bool = True - -# def tuple(self): -# return self.vae, self.source - - -# def is_automatic(): -# return shared.opts.sd_vae in {"Automatic", "auto"} # "auto" for people with old config - - -# def resolve_vae_from_setting() -> VaeResolution: -# if shared.opts.sd_vae == "None": -# return VaeResolution() - -# vae_from_options = vae_dict.get(shared.opts.sd_vae, None) -# if vae_from_options is not None: -# return VaeResolution(vae_from_options, "specified in settings") - -# if not is_automatic(): -# print(f"Couldn't find VAE named {shared.opts.sd_vae}; using None instead") - -# return VaeResolution(resolved=False) - - -# def resolve_vae_from_user_metadata(checkpoint_file) -> VaeResolution: -# metadata = extra_networks.get_user_metadata(checkpoint_file) -# vae_metadata = metadata.get("vae", None) -# if vae_metadata is not None and vae_metadata != "Automatic": -# if vae_metadata == "None": -# return VaeResolution() - -# vae_from_metadata = vae_dict.get(vae_metadata, None) -# if vae_from_metadata is not None: -# return VaeResolution(vae_from_metadata, "from user metadata") - -# return VaeResolution(resolved=False) - - -# def resolve_vae_near_checkpoint(checkpoint_file) -> VaeResolution: -# vae_near_checkpoint = find_vae_near_checkpoint(checkpoint_file) -# if vae_near_checkpoint is not None and (not shared.opts.sd_vae_overrides_per_model_preferences or is_automatic()): -# return VaeResolution(vae_near_checkpoint, "found near the checkpoint") - -# return VaeResolution(resolved=False) - - -# def resolve_vae(checkpoint_file) -> VaeResolution: -# if shared.cmd_opts.vae_path is not None: -# return VaeResolution(shared.cmd_opts.vae_path, "from commandline argument") - -# if shared.opts.sd_vae_overrides_per_model_preferences and not is_automatic(): -# return resolve_vae_from_setting() - -# res = resolve_vae_from_user_metadata(checkpoint_file) -# if res.resolved: -# return res - -# res = resolve_vae_near_checkpoint(checkpoint_file) -# if res.resolved: -# return res - -# res = resolve_vae_from_setting() - -# return res - - -def reload_vae_weights(vae: str): +def reload_vae_weights(vae: str) -> bool: if vae in (None, "None", "Automatic"): - return + return False store_base_vae(shared.sd_model) vae_sd = utils.load_torch_file(vae) _load_vae_dict(shared.sd_model, vae_sd) + return True def restore_vae_weights(): diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index ecb0bad9..3b6380e8 100644 --- a/scripts/xyz_grid.py +++ b/scripts/xyz_grid.py @@ -115,8 +115,8 @@ def apply_size(p: StableDiffusionProcessing, x: str, _): logger.error(f'Invalid Size "{x}" for X/Y/Z Plot') -def apply_vae(p: StableDiffusionProcessing, x: str, _): - p.override_settings["sd_vae"] = find_vae(x) +def apply_vae(p: StableDiffusionProcessing, x: str, xs: list[str]): + p.override_settings["sd_vae"] = (find_vae(x), xs) def apply_styles(p: StableDiffusionProcessing, x: str, _):