diff --git a/modules/sd_samplers_common.py b/modules/sd_samplers_common.py index ddad48d6..5e1b7238 100644 --- a/modules/sd_samplers_common.py +++ b/modules/sd_samplers_common.py @@ -1,5 +1,9 @@ import inspect from collections import namedtuple +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from backend.diffusion_engine.base import ForgeDiffusionEngine import k_diffusion.sampling import numpy as np @@ -246,13 +250,15 @@ def apply_refiner(cfg_denoiser, x, sigma): cfg_denoiser.p.extra_generation_params["Refiner switch at"] = refiner_switch_at if opts.refiner_fast_sd: + sd_model: "ForgeDiffusionEngine" = shared.sd_model + 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 + model = sd_model.forge_objects.unet.model.diffusion_model sd = load_torch_file(refiner_checkpoint_info.filename) sd = preprocess_state_dict(sd) @@ -265,9 +271,28 @@ def apply_refiner(cfg_denoiser, x, sigma): ORIGINAL_CHECKPOINT = shared.sd_model.sd_checkpoint_info.filename load_state_dict(model, sd) + # 1. reset the current_lora_hash so networks.py/load_networks() parse the LoRA again + sd_model.current_lora_hash = str([]) + + # 2. parse the LoRA to save to ModelPatcher.lora_patches + if not cfg_denoiser.p.disable_extra_networks: + loras = cfg_denoiser.p.extra_network_data.pop("lora", None) + cfg_denoiser.p.extra_network_data["lora"] = apply_lora_for_refiner(loras) + extra_networks.activate(cfg_denoiser.p, cfg_denoiser.p.extra_network_data) + + # 3. reset the loaded_hash so LoraLoader load the LoRA again + sd_model.forge_objects.unet.lora_loader.loaded_hash = str([]) + + # 4. actually load the LoRA + sd_model.forge_objects.unet.refresh_loras() + + # 5. reset the hashes again for the non-refiner pass + sd_model.current_lora_hash = str([]) + sd_model.forge_objects.unet.lora_loader.loaded_hash = str([]) + return True - sampling_cleanup(sd_models.model_data.get_sd_model().forge_objects.unet) + sampling_cleanup(shared.sd_model.forge_objects.unet) original_checkpoint = getattr(shared.opts, "sd_model_checkpoint") checkpoint_changed = main_entry.checkpoint_change(refiner_checkpoint_info.short_title, preset=None, save=False, refresh=False) @@ -290,7 +315,7 @@ def apply_refiner(cfg_denoiser, x, sigma): cfg_denoiser.p.setup_conds() cfg_denoiser.update_inner_model() - sampling_prepare(sd_models.model_data.get_sd_model().forge_objects.unet, x=x) + sampling_prepare(shared.sd_model.forge_objects.unet, x=x) return True diff --git a/modules/shared_options.py b/modules/shared_options.py index b0da9fba..0117a789 100644 --- a/modules/shared_options.py +++ b/modules/shared_options.py @@ -369,7 +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_fast_sd": OptionInfo(False, 'Reload "state_dict" Only').info("EXPERIMENTAL"), "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",