This commit is contained in:
Haoming 2025-10-29 11:14:50 +08:00
parent 746c116004
commit f9f0954024
5 changed files with 28 additions and 177 deletions

View File

@ -233,7 +233,6 @@ class Api:
self.add_api_route("/sdapi/v1/refresh-vae", self.refresh_vae, methods=["POST"])
self.add_api_route("/sdapi/v1/memory", self.get_memory, methods=["GET"], response_model=models.MemoryResponse)
self.add_api_route("/sdapi/v1/unload-checkpoint", self.unloadapi, methods=["POST"])
self.add_api_route("/sdapi/v1/reload-checkpoint", self.reloadapi, methods=["POST"])
self.add_api_route("/sdapi/v1/scripts", self.get_scripts_list, methods=["GET"], response_model=models.ScriptsList)
self.add_api_route("/sdapi/v1/script-info", self.get_script_info, methods=["GET"], response_model=list[models.ScriptInfo])
self.add_api_route("/sdapi/v1/extensions", self.get_extensions_list, methods=["GET"], response_model=list[models.ExtensionItem])
@ -649,11 +648,6 @@ class Api:
return {}
def reloadapi(self):
sd_models.send_model_to_device(shared.sd_model)
return {}
def skip(self):
shared.state.skip()

View File

@ -170,18 +170,13 @@ def configure_sigint_handler():
def configure_opts_onchange():
from modules import shared, sd_models, sd_vae, ui_tempdir
from modules import shared, sd_vae, ui_tempdir
from modules.call_queue import wrap_queued_call
from modules_forge import main_thread
# shared.opts.onchange("sd_model_checkpoint", wrap_queued_call(lambda: main_thread.run_and_wait_result(sd_models.reload_model_weights)), call=False)
# shared.opts.onchange("sd_vae", wrap_queued_call(lambda: main_thread.run_and_wait_result(sd_vae.reload_vae_weights)), call=False)
shared.opts.onchange("sd_vae_overrides_per_model_preferences", wrap_queued_call(lambda: main_thread.run_and_wait_result(sd_vae.reload_vae_weights)), call=False)
shared.opts.onchange("temp_dir", ui_tempdir.on_tmpdir_changed)
shared.opts.onchange("gradio_theme", shared.reload_gradio_theme)
# shared.opts.onchange("cross_attention_optimization", wrap_queued_call(lambda: sd_hijack.model_hijack.redo_hijack(shared.sd_model)), call=False)
# shared.opts.onchange("fp8_storage", wrap_queued_call(lambda: sd_models.reload_model_weights()), call=False)
# shared.opts.onchange("cache_fp16_weight", wrap_queued_call(lambda: sd_models.reload_model_weights(forced_reload=True)), call=False)
startup_timer.record("opts onchange")

View File

@ -978,9 +978,6 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
if p.n_iter > 1:
shared.state.job = f"Batch {n+1} out of {p.n_iter}"
# TODO: This currently seems broken. It should be fixed or removed.
sd_models.apply_alpha_schedule_override(p.sd_model, p)
sigmas_backup = None
if (opts.sd_noise_schedule == "Zero Terminal SNR" or getattr(p.sd_model.model_config, "ztsnr", False)) and p is not None:
p.extra_generation_params["Noise Schedule"] = "Zero Terminal SNR"
@ -1409,33 +1406,32 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
else:
decoded_samples = None
with sd_models.SkipWritingToConfig():
fp_checkpoint = getattr(shared.opts, "sd_model_checkpoint")
fp_additional_modules = getattr(shared.opts, "forge_additional_modules")
fp_checkpoint = getattr(shared.opts, "sd_model_checkpoint")
fp_additional_modules = getattr(shared.opts, "forge_additional_modules")
reload = False
if hasattr(self, "hr_additional_modules") and "Use same choices" not in self.hr_additional_modules:
modules_changed = main_entry.modules_change(self.hr_additional_modules, preset=None, save=False, refresh=False)
if modules_changed:
reload = True
reload = False
if hasattr(self, "hr_additional_modules") and "Use same choices" not in self.hr_additional_modules:
modules_changed = main_entry.modules_change(self.hr_additional_modules, preset=None, save=False, refresh=False)
if modules_changed:
reload = True
if self.hr_checkpoint_name and self.hr_checkpoint_name != "Use same checkpoint":
checkpoint_changed = main_entry.checkpoint_change(self.hr_checkpoint_name, preset=None, save=False, refresh=False)
if checkpoint_changed:
self.firstpass_use_distilled_cfg_scale = self.sd_model.use_distilled_cfg_scale
reload = True
if self.hr_checkpoint_name and self.hr_checkpoint_name != "Use same checkpoint":
checkpoint_changed = main_entry.checkpoint_change(self.hr_checkpoint_name, preset=None, save=False, refresh=False)
if checkpoint_changed:
self.firstpass_use_distilled_cfg_scale = self.sd_model.use_distilled_cfg_scale
reload = True
if reload:
try:
main_entry.refresh_model_loading_parameters()
sd_models.forge_model_reload()
finally:
main_entry.modules_change(fp_additional_modules, preset=None, save=False, refresh=False)
main_entry.checkpoint_change(fp_checkpoint, preset=None, save=False, refresh=False)
main_entry.refresh_model_loading_parameters()
if reload:
try:
main_entry.refresh_model_loading_parameters()
sd_models.forge_model_reload()
finally:
main_entry.modules_change(fp_additional_modules, preset=None, save=False, refresh=False)
main_entry.checkpoint_change(fp_checkpoint, preset=None, save=False, refresh=False)
main_entry.refresh_model_loading_parameters()
if self.sd_model.use_distilled_cfg_scale:
self.extra_generation_params["Hires Distilled CFG Scale"] = self.hr_distilled_cfg
if self.sd_model.use_distilled_cfg_scale:
self.extra_generation_params["Hires Distilled CFG Scale"] = self.hr_distilled_cfg
return self.sample_hr_pass(samples, decoded_samples, seeds, subseeds, subseed_strength, prompts)

View File

@ -1,5 +1,3 @@
import contextlib
import enum
import gc
import math
import os
@ -20,7 +18,6 @@ model_path = os.path.abspath(os.path.join(paths.models_path, model_dir))
checkpoints_list = {}
checkpoint_aliases = {}
checkpoint_alisases = checkpoint_aliases # for compatibility with old name
def replace_key(d, key, new_key, value):
@ -121,13 +118,8 @@ class CheckpointInfo:
def setup_model():
"""called once at startup to do various one-time tasks related to SD models"""
os.makedirs(model_path, exist_ok=True)
enable_midas_autodownload()
patch_given_betas()
def checkpoint_tiles(use_short=False):
return [x.short_title if use_short else x.name for x in checkpoints_list.values()]
@ -215,14 +207,6 @@ def select_checkpoint():
return checkpoint_info
def transform_checkpoint_dict_key(k, replacements):
pass
def get_state_dict_from_checkpoint(pl_sd):
pass
def read_metadata_from_safetensors(filename):
import json
@ -251,86 +235,9 @@ def read_metadata_from_safetensors(filename):
return res
def read_state_dict(checkpoint_file, print_global_state=False, map_location=None):
pass
def SkipWritingToConfig():
return contextlib.nullcontext()
def check_fp8(model):
pass
def set_model_type(model, state_dict):
pass
def set_model_fields(model):
pass
def load_model_weights(model, checkpoint_info: CheckpointInfo, state_dict, timer):
pass
def enable_midas_autodownload():
pass
def patch_given_betas():
pass
def repair_config(sd_config, state_dict=None):
pass
def rescale_zero_terminal_snr_abar(alphas_cumprod):
alphas_bar_sqrt = alphas_cumprod.sqrt()
# Store old values.
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
# Shift so the last timestep is zero.
alphas_bar_sqrt -= alphas_bar_sqrt_T
# Scale so the first timestep is back to the old value.
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
# Convert alphas_bar_sqrt to betas
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
alphas_bar[-1] = 4.8973451890853435e-08
return alphas_bar
def apply_alpha_schedule_override(sd_model, p=None):
"""
Applies an override to the alpha schedule of the model according to settings.
- downcasts the alpha schedule to half precision
- rescales the alpha schedule to have zero terminal SNR
"""
if not hasattr(sd_model, "alphas_cumprod") or not hasattr(sd_model, "alphas_cumprod_original"):
return
sd_model.alphas_cumprod = sd_model.alphas_cumprod_original.to(shared.device)
if opts.use_downcasted_alpha_bar:
if p is not None:
p.extra_generation_params["Downcast alphas_cumprod"] = opts.use_downcasted_alpha_bar
sd_model.alphas_cumprod = sd_model.alphas_cumprod.half().to(shared.device)
if opts.sd_noise_schedule == "Zero Terminal SNR":
if p is not None:
p.extra_generation_params["Noise Schedule"] = opts.sd_noise_schedule
sd_model.alphas_cumprod = rescale_zero_terminal_snr_abar(sd_model.alphas_cumprod).to(shared.device)
# This is a dummy class for backward compatibility when model is not load - for extensions like prompt all in one.
class FakeInitialModel:
"""a dummy class for compatibility when no model is loaded yet"""
def __init__(self):
self.cond_stage_model = None
self.chunk_length = 75
@ -356,46 +263,6 @@ class SdModelData:
model_data = SdModelData()
def get_empty_cond(sd_model):
pass
def send_model_to_cpu(m):
pass
def model_target_device(m):
return devices.device
def send_model_to_device(m):
pass
def send_model_to_trash(m):
pass
def instantiate_from_config(config, state_dict=None):
pass
def get_obj_from_str(string, reload=False):
pass
def load_model(checkpoint_info=None, already_loaded_state_dict=None):
pass
def reuse_model_from_already_loaded(sd_model, checkpoint_info, timer):
pass
def reload_model_weights(sd_model=None, info=None, forced_reload=False):
pass
def unload_model_weights(sd_model=None, info=None):
memory_management.unload_all_models()
return

View File

@ -419,13 +419,12 @@ def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend
return processed_result
class SharedSettingsStackHelper(object):
class SharedSettingsStackHelper:
def __enter__(self):
pass
def __exit__(self, exc_type, exc_value, tb):
modules.sd_models.reload_model_weights()
modules.sd_vae.reload_vae_weights()
def __exit__(self, *args, **kwargs):
pass
re_range = re.compile(r"\s*([+-]?\s*\d+)\s*-\s*([+-]?\s*\d+)(?:\s*\(([+-]\d+)\s*\))?\s*")