fast refiner

This commit is contained in:
Haoming 2026-02-20 19:39:07 +08:00
parent 2109c0e9f2
commit edba28dacc
3 changed files with 58 additions and 0 deletions

View File

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

View File

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

View File

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