This commit is contained in:
Haoming 2026-04-28 13:06:57 +08:00
parent ede717bbf7
commit de1eff76bb
3 changed files with 20 additions and 91 deletions

View File

@ -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

View File

@ -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():

View File

@ -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, _):