mirror of
https://github.com/lllyasviel/stable-diffusion-webui-forge.git
synced 2026-07-21 21:01:24 +08:00
vae
This commit is contained in:
parent
ede717bbf7
commit
de1eff76bb
@ -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
|
||||
|
||||
|
||||
@ -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():
|
||||
|
||||
@ -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, _):
|
||||
|
||||
Loading…
Reference in New Issue
Block a user