mirror of
https://github.com/lllyasviel/stable-diffusion-webui-forge.git
synced 2026-07-21 21:01:24 +08:00
fast refiner
This commit is contained in:
parent
2109c0e9f2
commit
edba28dacc
@ -1,4 +1,5 @@
|
||||
import gradio as gr
|
||||
import torch
|
||||
|
||||
from modules import scripts, sd_models
|
||||
from modules.infotext_utils import PasteField
|
||||
@ -53,3 +54,30 @@ class ScriptRefiner(scripts.ScriptBuiltinUI):
|
||||
else:
|
||||
p.refiner_checkpoint = refiner_checkpoint
|
||||
p.refiner_switch_at = refiner_switch_at
|
||||
|
||||
@torch.inference_mode()
|
||||
def postprocess(self, *args, **kwargs):
|
||||
from modules import sd_samplers_common
|
||||
|
||||
if sd_samplers_common.ORIGINAL_CHECKPOINT is None:
|
||||
return
|
||||
|
||||
import huggingface_guess
|
||||
|
||||
from backend.loader import preprocess_state_dict
|
||||
from backend.state_dict import load_state_dict, try_filter_state_dict
|
||||
from backend.utils import load_torch_file
|
||||
from modules_forge.main_entry import logger
|
||||
|
||||
model = sd_models.model_data.get_sd_model().forge_objects.unet.model.diffusion_model
|
||||
|
||||
sd = load_torch_file(sd_samplers_common.ORIGINAL_CHECKPOINT)
|
||||
sd = preprocess_state_dict(sd)
|
||||
|
||||
guess = huggingface_guess.guess(sd)
|
||||
|
||||
sd = try_filter_state_dict(sd, guess.unet_key_prefix)
|
||||
|
||||
logger.info("Restoring state_dict...")
|
||||
sd_samplers_common.ORIGINAL_CHECKPOINT = None
|
||||
load_state_dict(model, sd)
|
||||
|
||||
@ -185,6 +185,9 @@ def apply_lora_for_refiner(loras: list[extra_networks.ExtraNetworkParams]):
|
||||
return result
|
||||
|
||||
|
||||
ORIGINAL_CHECKPOINT: str = None
|
||||
|
||||
|
||||
def apply_refiner(cfg_denoiser, x, sigma):
|
||||
if not (refiner_switch_at := cfg_denoiser.p.refiner_switch_at):
|
||||
return False
|
||||
@ -196,6 +199,10 @@ def apply_refiner(cfg_denoiser, x, sigma):
|
||||
if float(sigma) > refiner_switch_at:
|
||||
return False
|
||||
|
||||
global ORIGINAL_CHECKPOINT
|
||||
if ORIGINAL_CHECKPOINT is not None:
|
||||
return False
|
||||
|
||||
refiner_checkpoint_info = cfg_denoiser.p.refiner_checkpoint_info
|
||||
if refiner_checkpoint_info is None or shared.sd_model.sd_checkpoint_info == refiner_checkpoint_info:
|
||||
return False
|
||||
@ -207,6 +214,28 @@ def apply_refiner(cfg_denoiser, x, sigma):
|
||||
cfg_denoiser.p.extra_generation_params["Refiner"] = refiner_checkpoint_info.short_title
|
||||
cfg_denoiser.p.extra_generation_params["Refiner switch at"] = refiner_switch_at
|
||||
|
||||
if opts.refiner_fast_sd:
|
||||
import huggingface_guess
|
||||
|
||||
from backend.loader import preprocess_state_dict
|
||||
from backend.state_dict import load_state_dict, try_filter_state_dict
|
||||
from backend.utils import load_torch_file
|
||||
|
||||
model = sd_models.model_data.get_sd_model().forge_objects.unet.model.diffusion_model
|
||||
|
||||
sd = load_torch_file(refiner_checkpoint_info.filename)
|
||||
sd = preprocess_state_dict(sd)
|
||||
|
||||
guess = huggingface_guess.guess(sd)
|
||||
|
||||
sd = try_filter_state_dict(sd, guess.unet_key_prefix)
|
||||
|
||||
main_entry.logger.info("Reloading state_dict...")
|
||||
ORIGINAL_CHECKPOINT = shared.sd_model.sd_checkpoint_info.filename
|
||||
load_state_dict(model, sd)
|
||||
|
||||
return True
|
||||
|
||||
sampling_cleanup(sd_models.model_data.get_sd_model().forge_objects.unet)
|
||||
|
||||
original_checkpoint = getattr(shared.opts, "sd_model_checkpoint")
|
||||
|
||||
@ -369,6 +369,7 @@ options_templates.update(
|
||||
("refiner", "Refiner", "sd"),
|
||||
{
|
||||
"show_refiner": OptionInfo(False, "Display the Refiner Accordion").info("Refiner swaps the model in the middle of generation; useful for Wan 2.2 <b>High Noise</b> to <b>Low Noise</b> switching").needs_reload_ui(),
|
||||
"refiner_fast_sd": OptionInfo(False, 'Reload "state_dict" Only').info("EXPERIMENTAL").info("does not support LoRA currently"),
|
||||
"refiner_use_steps": OptionInfo(False, 'Switch based on "steps" instead').info('by default, Refiner swaps the model based on "sigmas" to match <a href="https://www.reddit.com/r/StableDiffusion/comments/1n3qns1/wan_22_how_many_highsteps_are_needed_a_simple/">Wan 2.2</a> \'s behavior'),
|
||||
"refiner_lora_replacement": OptionInfo(
|
||||
"high_noise=low_noise",
|
||||
|
||||
Loading…
Reference in New Issue
Block a user