Support kolors for custom guider
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
|
||||
+1
-2
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user