diff --git a/py/kolors/model_patch.py b/py/kolors/model_patch.py index 3b2ca1d..02da46e 100644 --- a/py/kolors/model_patch.py +++ b/py/kolors/model_patch.py @@ -1,9 +1,15 @@ import torch.nn import comfy.model_management import comfy.samplers +from comfy_extras.nodes_custom_sampler import Guider_Basic, Guider_DualCFG -if "original_CFGGuider_inner_set_conds" not in globals(): +_globals = globals() +if "original_CFGGuider_inner_set_conds" not in _globals: original_CFGGuider_inner_set_conds = comfy.samplers.CFGGuider.set_conds +if "original_BasicGuider_inner_set_conds" not in _globals: + original_BasicGuider_inner_set_conds = Guider_Basic.set_conds +if "original_DualCFGGuider_inner_set_conds" not in _globals: + original_DualCFGGuider_inner_set_conds = Guider_DualCFG.set_conds def add_model_patch(model, sd): load_device = comfy.model_management.get_torch_device() @@ -17,7 +23,7 @@ def add_model_patch(model, sd): "hid_proj": hid_proj } -def patched_kolors_conds(self, positive, negative): +def patched_set_conds(self, positive, negative=None, middle=None): if "model_patch" in self.model_options: mp = self.model_options["model_patch"] if "hid_proj" in mp: @@ -33,14 +39,33 @@ def patched_kolors_conds(self, positive, negative): if hasattr(positive[i][1]["control"], "control_model"): positive[i][1]["control"].control_model.label_emb = self.model_patcher.model.diffusion_model.label_emb - for i in range(len(negative)): - negative[i][0] = hid_proj(negative[i][0]) - if "control" in negative[i][1]: - if hasattr(negative[i][1]["control"], "control_model"): - negative[i][1]["control"].control_model.label_emb = self.model_patcher.model.diffusion_model.label_emb + if negative is not None: + for i in range(len(negative)): + negative[i][0] = hid_proj(negative[i][0]) + if "control" in negative[i][1]: + if hasattr(negative[i][1]["control"], "control_model"): + negative[i][1]["control"].control_model.label_emb = self.model_patcher.model.diffusion_model.label_emb + if middle is not None: + for i in range(len(middle)): + middle[i][0] = hid_proj(middle[i][0]) + if "control" in middle[i][1]: + if hasattr(middle[i][1]["control"], "control_model"): + middle[i][1]["control"].control_model.label_emb = self.model_patcher.model.diffusion_model.label_emb + return self, positive, negative, middle + +def patched_cfgguider_set_conds(self, positive, negative): + self, positive, negative, _ = patched_set_conds(self, positive, negative) return original_CFGGuider_inner_set_conds(self, positive, negative) +def patched_basicguider_set_conds(self, positive): + self, positive, _, _ = patched_set_conds(self, positive) + return original_BasicGuider_inner_set_conds(self, positive) +def patched_dualcfgguider_set_conds(self, positive, midele, negative): + self, positive, negative, middle = patched_set_conds(self, positive, midele, negative) + return original_DualCFGGuider_inner_set_conds(self, positive, midele, negative) -comfy.samplers.CFGGuider.set_conds = patched_kolors_conds +comfy.samplers.CFGGuider.set_conds = patched_cfgguider_set_conds +Guider_Basic.set_conds = patched_basicguider_set_conds +Guider_DualCFG.set_conds = patched_dualcfgguider_set_conds diff --git a/py/libs/sampler.py b/py/libs/sampler.py index 4a7572c..3163955 100644 --- a/py/libs/sampler.py +++ b/py/libs/sampler.py @@ -8,7 +8,6 @@ from nodes import MAX_RESOLUTION from PIL import Image from typing import Dict, List, Optional, Tuple, Union, Any from ..brushnet.model_patch import add_model_patch -from ..kolors.model_patch import patched_kolors_conds class easySampler: def __init__(self): @@ -161,7 +160,7 @@ class easySampler: # noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, # seed=seed) - model, positive, negative = patched_kolors_conds(model, positive, negative) + # model, positive, negative = patched_kolors_conds(model, positive, negative) samples = comfy.samplers.sample(model, noise, positive, negative, cfg, device, _sampler, sigmas, latent_image=latent_image, model_options=model.model_options, denoise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=seed)