From edba28daccc56c9fc5eb6e2230de233a4a93e11f Mon Sep 17 00:00:00 2001 From: Haoming Date: Fri, 20 Feb 2026 19:39:07 +0800 Subject: [PATCH] fast refiner --- modules/processing_scripts/refiner.py | 28 ++++++++++++++++++++++++++ modules/sd_samplers_common.py | 29 +++++++++++++++++++++++++++ modules/shared_options.py | 1 + 3 files changed, 58 insertions(+) diff --git a/modules/processing_scripts/refiner.py b/modules/processing_scripts/refiner.py index 4b562da7..76062448 100644 --- a/modules/processing_scripts/refiner.py +++ b/modules/processing_scripts/refiner.py @@ -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) diff --git a/modules/sd_samplers_common.py b/modules/sd_samplers_common.py index 413a3f4d..923265a3 100644 --- a/modules/sd_samplers_common.py +++ b/modules/sd_samplers_common.py @@ -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") diff --git a/modules/shared_options.py b/modules/shared_options.py index b73a81ae..b0da9fba 100644 --- a/modules/shared_options.py +++ b/modules/shared_options.py @@ -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 High Noise to Low Noise 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 Wan 2.2 \'s behavior'), "refiner_lora_replacement": OptionInfo( "high_noise=low_noise",