refiner lora

This commit is contained in:
Haoming 2026-02-26 12:20:12 +08:00
parent d28b63ad9b
commit aa07cb5fed
2 changed files with 29 additions and 4 deletions

View File

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

View File

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