Support kolors for custom guider

This commit is contained in:
yolain
2024-07-12 19:09:25 +08:00
parent 1eb1e1a5c1
commit ee09de3f16
2 changed files with 34 additions and 10 deletions
+33 -8
View File
@@ -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
View File
@@ -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)