diff --git a/backend/diffusion_engine/base.py b/backend/diffusion_engine/base.py index 15dc092f..55d24cc0 100644 --- a/backend/diffusion_engine/base.py +++ b/backend/diffusion_engine/base.py @@ -3,14 +3,18 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: import torch + from backend.patcher.clip import CLIP + from backend.patcher.unet import UnetPatcher + from backend.patcher.vae import VAE + from backend import memory_management, utils class ForgeObjects: def __init__(self, unet, clip, vae, clipvision): - self.unet = unet - self.clip = clip - self.vae = vae + self.unet: "UnetPatcher" = unet + self.clip: "CLIP" = clip + self.vae: "VAE" = vae self.clipvision = clipvision def shallow_copy(self): @@ -24,9 +28,9 @@ class ForgeDiffusionEngine: self.model_config = estimated_config self.is_inpaint = estimated_config.inpaint_model() - self.forge_objects = None - self.forge_objects_original = None - self.forge_objects_after_applying_lora = None + self.forge_objects: "ForgeObjects" = None + self.forge_objects_original: "ForgeObjects" = None + self.forge_objects_after_applying_lora: "ForgeObjects" = None self.current_lora_hash = str([]) @@ -58,8 +62,6 @@ class ForgeDiffusionEngine: def fix_for_webui_backward_compatibility(self): self.tiling_enabled = False - self.first_stage_model = None - self.cond_stage_model = None self.use_distilled_cfg_scale = False self.use_shift = False self.is_sd1 = False @@ -67,6 +69,20 @@ class ForgeDiffusionEngine: self.is_flux = False # affects the usage of TAESD self.is_wan = False # affects the usage of WanVAE (B, C, T, H, W) + @property + def first_stage_model(self): + try: + return self.forge_objects.vae.first_stage_model + except Exception: + return None + + @property + def cond_stage_model(self): + try: + return self.forge_objects.clip.cond_stage_model + except Exception: + return None + def clear_references(self): # called by ImageStitch self.ref_latents.clear() diff --git a/modules/processing.py b/modules/processing.py index ea554a43..50b500f4 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -817,6 +817,8 @@ 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) + # apply any options overrides set_config(p.override_settings, is_api=True, run_callbacks=False, save_config=False) @@ -826,6 +828,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed: pass else: manage_model_and_prompt_cache(p) + sd_vae.reload_vae_weights(_vae_override) # backwards compatibility, fix sampler and scheduler if invalid sd_samplers.fix_p_invalid_sampler_and_scheduler(p) @@ -837,6 +840,8 @@ def process_images(p: StableDiffusionProcessing) -> Processed: # restore original options if p.override_settings_restore_afterwards: set_config(stored_opts, save_config=False) + if _vae_override is not None: + sd_vae.restore_vae_weights() return res diff --git a/modules/sd_vae.py b/modules/sd_vae.py index dc37d144..504c33d0 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -1,63 +1,73 @@ import glob -import os +import os.path from copy import deepcopy -from dataclasses import dataclass -from modules import extra_networks, hashes, paths, sd_models, shared +import torch + +from backend import memory_management, utils +from modules import hashes, paths, sd_models, shared vae_path = os.path.abspath(os.path.join(paths.models_path, "VAE")) -vae_dict = {} +vae_ignore_keys: set[str] = {"model_ema.decay", "model_ema.num_updates"} +vae_dict: dict[str, os.PathLike] = {} -base_vae = None -loaded_vae_file = None -checkpoint_info = None +base_vae: dict[str, torch.Tensor] = None +loaded_vae_file: os.PathLike = None +checkpoint_info: "sd_models.CheckpointInfo" = None -def get_loaded_vae_name(): +@torch.inference_mode() +def _load_vae_dict(model, vae_sd: dict): + sd = {k: v for k, v in vae_sd.items() if k[0:4] != "loss" and k not in vae_ignore_keys} + model.first_stage_model.load_state_dict(sd) + + +def get_loaded_vae_name() -> str: if loaded_vae_file is None: return None - return os.path.basename(loaded_vae_file) -def get_loaded_vae_hash(): +def get_loaded_vae_hash() -> str: if loaded_vae_file is None: return None sha256 = hashes.sha256(loaded_vae_file, "vae") - 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 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 - if checkpoint_info != model.sd_checkpoint_info: - assert not loaded_vae_file, "Trying to store non-base VAE!" - base_vae = deepcopy(model.first_stage_model.state_dict()) - checkpoint_info = model.sd_checkpoint_info + assert loaded_vae_file is None + memory_management.logger.debug("Storing Original VAE...") + base_vae = deepcopy(model.first_stage_model.state_dict()) + checkpoint_info = model.sd_checkpoint_info def delete_base_vae(): global base_vae, checkpoint_info base_vae = None checkpoint_info = None + memory_management.soft_empty_cache() def restore_base_vae(model): global loaded_vae_file - if base_vae is not None and checkpoint_info == model.sd_checkpoint_info: - print("Restoring base VAE") - loaded_vae_file = None + if base_vae is None: + return + memory_management.logger.debug("Restoring Original VAE...") + _load_vae_dict(model, base_vae) + loaded_vae_file = None delete_base_vae() -def get_filename(filepath): +def get_filename(filepath: os.PathLike) -> str: return os.path.basename(filepath) @@ -86,80 +96,93 @@ 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 +# 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 +# return None -@dataclass -class VaeResolution: - vae: str = None - source: str = None - resolved: bool = True +# @dataclass +# class VaeResolution: +# vae: str = None +# source: str = None +# resolved: bool = True - def tuple(self): - return self.vae, self.source +# 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 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() +# 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") +# 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") +# if not is_automatic(): +# print(f"Couldn't find VAE named {shared.opts.sd_vae}; using None instead") - return VaeResolution(resolved=False) +# 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() +# 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") +# 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) +# 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") +# 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) +# 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") +# 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() +# 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_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_near_checkpoint(checkpoint_file) +# if res.resolved: +# return res - res = resolve_vae_from_setting() +# res = resolve_vae_from_setting() - return res +# return res + + +def reload_vae_weights(vae: str): + if vae in (None, "None", "Automatic"): + return + + store_base_vae(shared.sd_model) + vae_sd = utils.load_torch_file(vae) + _load_vae_dict(shared.sd_model, vae_sd) + + +def restore_vae_weights(): + restore_base_vae(shared.sd_model) diff --git a/modules/shared_options.py b/modules/shared_options.py index a090732f..7da63388 100644 --- a/modules/shared_options.py +++ b/modules/shared_options.py @@ -260,7 +260,7 @@ image to and from latent space representation. Latent space is what Stable Diffu to create the resulting image after the sampling is finished. For img2img, VAE is additionally used to process user's input image before the sampling. """ ), - "sd_vae": OptionInfo("Automatic", "SD VAE", gr.Dropdown, lambda: {"choices": shared_items.sd_vae_items()}, refresh=shared_items.refresh_vae_list, infotext="VAE").info("None = always use VAE from checkpoint; Automatic = use VAE with the same filename as checkpoint"), + "sd_vae": OptionInfo("Automatic", "SD VAE", gr.Dropdown, {"choices": ("Automatic",), "interactive": False}), "sd_vae_overrides_per_model_preferences": OptionInfo(True, '"SD VAE" option overrides per-model preference'), "sd_vae_encode_method": OptionInfo("Full", "VAE for Encoding", gr.Radio, {"choices": ("Full", "TAESD")}, infotext="VAE Encoder").info("method to encode image to latent (img2img / Hires. fix / inpaint)"), "sd_vae_decode_method": OptionInfo("Full", "VAE for Decoding", gr.Radio, {"choices": ("Full", "TAESD")}, infotext="VAE Decoder").info("method to decode latent to image"), diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index f5b7f70f..5ba39bef 100644 --- a/scripts/xyz_grid.py +++ b/scripts/xyz_grid.py @@ -1,26 +1,35 @@ -from collections import namedtuple -from copy import copy -from itertools import permutations, chain -import random import csv import os.path +import random +import re +from collections import namedtuple +from copy import copy from io import StringIO -from PIL import Image +from itertools import chain, permutations + +import gradio as gr import numpy as np +from PIL import Image import modules.scripts as scripts -import gradio as gr - -from modules import images, sd_samplers, processing, sd_models, sd_vae, sd_schedulers, errors -from modules.processing import process_images, Processed, StableDiffusionProcessingTxt2Img -from modules.shared import opts, state -from modules.sd_models import model_data, select_checkpoint -import modules.shared as shared -import modules.sd_samplers import modules.sd_models -import modules.sd_vae -import re - +import modules.shared as shared +from modules import ( + errors, + images, + processing, + sd_models, + sd_samplers, + sd_schedulers, + sd_vae, +) +from modules.processing import ( + Processed, + StableDiffusionProcessingTxt2Img, + process_images, +) +from modules.sd_models import model_data, select_checkpoint +from modules.shared import opts, state from modules.ui_components import ToolButton fill_values_symbol = "\U0001f4d2" # 📒 @@ -139,16 +148,17 @@ def apply_size(p, x: str, xs) -> None: print(f"Invalid size in XYZ plot: {x}") -def find_vae(name: str): - if (name := name.strip().lower()) in ('auto', 'automatic'): - return 'Automatic' - elif name == 'none': - return 'None' - return next((k for k in modules.sd_vae.vae_dict if k.lower() == name), print(f'No VAE found for {name}; using Automatic') or 'Automatic') +def find_vae(name: str) -> str: + if name is None or (name := name.strip().lower()) == "none": + return "None" + elif name in ("auto", "automatic"): + return "Automatic" + else: + return sd_vae.vae_dict[name] def apply_vae(p, x, xs): - p.override_settings['sd_vae'] = find_vae(x) + p.override_settings["sd_vae"] = find_vae(x) def apply_styles(p: StableDiffusionProcessingTxt2Img, x: str, _):