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
cb3c1fc535
commit
e9e2bd7339
@ -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()
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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"),
|
||||
|
||||
@ -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, _):
|
||||
|
||||
Loading…
Reference in New Issue
Block a user