rework refiner for some potential new features (#642)

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync

* sync
This commit is contained in:
lllyasviel
2023-10-11 03:07:43 -07:00
committed by GitHub
parent 777510dc9b
commit bbdf4bd120
10 changed files with 209 additions and 93 deletions
+9 -45
View File
@@ -4,7 +4,6 @@ patch_all()
import os
import random
import einops
import torch
import numpy as np
@@ -63,47 +62,6 @@ def apply_controlnet(positive, negative, control_net, image, strength, start_per
image=image, strength=strength, start_percent=start_percent, end_percent=end_percent)
@torch.no_grad()
@torch.inference_mode()
def load_unet_only(unet_path):
sd_raw = comfy.utils.load_torch_file(unet_path)
sd = {}
flag = 'model.diffusion_model.'
for k in list(sd_raw.keys()):
if k.startswith(flag):
sd[k[len(flag):]] = sd_raw[k]
del sd_raw[k]
parameters = comfy.utils.calculate_parameters(sd)
fp16 = comfy.model_management.should_use_fp16(model_params=parameters)
if "input_blocks.0.0.weight" in sd:
# ldm
model_config = comfy.model_detection.model_config_from_unet(sd, "", fp16)
if model_config is None:
raise RuntimeError("ERROR: Could not detect model type of: {}".format(unet_path))
new_sd = sd
else:
# diffusers
model_config = comfy.model_detection.model_config_from_diffusers_unet(sd, fp16)
if model_config is None:
print("ERROR UNSUPPORTED UNET", unet_path)
return None
diffusers_keys = comfy.utils.unet_to_diffusers(model_config.unet_config)
new_sd = {}
for k in diffusers_keys:
if k in sd:
new_sd[diffusers_keys[k]] = sd.pop(k)
else:
print(diffusers_keys[k], k)
offload_device = comfy.model_management.unet_offload_device()
model = model_config.get_model(new_sd, "")
model = model.to(offload_device)
model.load_model_weights(new_sd, "")
return comfy.model_patcher.ModelPatcher(model, load_device=comfy.model_management.get_torch_device(), offload_device=offload_device)
@torch.no_grad()
@torch.inference_mode()
def load_model(ckpt_filename):
@@ -239,7 +197,7 @@ def get_previewer():
@torch.inference_mode()
def ksampler(model, positive, negative, latent, seed=None, steps=30, cfg=7.0, sampler_name='dpmpp_fooocus_2m_sde_inpaint_seamless',
scheduler='karras', denoise=1.0, disable_noise=False, start_step=None, last_step=None,
force_full_denoise=False, callback_function=None, refiner=None, refiner_switch=-1):
force_full_denoise=False, callback_function=None, refiner=None, refiner_switch=-1, previewer_start=None, previewer_end=None):
latent_image = latent["samples"]
if disable_noise:
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
@@ -253,13 +211,19 @@ def ksampler(model, positive, negative, latent, seed=None, steps=30, cfg=7.0, sa
previewer = get_previewer()
if previewer_start is None:
previewer_start = 0
if previewer_end is None:
previewer_end = steps
def callback(step, x0, x, total_steps):
comfy.model_management.throw_exception_if_processing_interrupted()
y = None
if previewer is not None:
y = previewer(x0, step, total_steps)
y = previewer(x0, previewer_start + step, previewer_end)
if callback_function is not None:
callback_function(step, x0, x, total_steps, y)
callback_function(previewer_start + step, x0, x, previewer_end, y)
disable_pbar = False
modules.sample_hijack.current_refiner = refiner