This commit is contained in:
Haoming 2026-02-12 15:54:49 +08:00
parent cb3c1fc535
commit e9e2bd7339
5 changed files with 161 additions and 107 deletions

View File

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

View File

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

View File

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

View File

@ -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"),

View File

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