Fix some kolors logic

This commit is contained in:
yolain
2024-07-13 22:48:43 +08:00
parent 56c8b64bd1
commit c68258304c
5 changed files with 28 additions and 37 deletions
+14 -1
View File
@@ -4111,6 +4111,7 @@ class samplerSettingsNoiseIn:
# 预采样设置(自定义)
import comfy_extras.nodes_custom_sampler as custom_samplers
from .kolors.model_patch import patched_set_conds
from tqdm import trange
class samplerCustomSettings:
@@ -4320,10 +4321,22 @@ class samplerCustomSettings:
# guider
if guider == 'CFG':
# --------------------------------
# kolors patch set conditioning
model, positive, negative, _ = patched_set_conds(model, positive, negative, None)
# --------------------------------
_guider, = self.get_custom_cls('CFGGuider').get_guider(model, positive, negative, cfg)
elif guider in ['DualCFG', 'IP2P+DualCFG']:
_guider, = self.get_custom_cls('DualCFGGuider').get_guider(model, positive, negative, pipe['negative'], cfg, cfg_negative)
# --------------------------------
# kolors patch set conditioning
model, positive, negative, middle = patched_set_conds(model, positive, pipe['negative'], negative)
# --------------------------------
_guider, = self.get_custom_cls('DualCFGGuider').get_guider(model, positive, middle, negative, cfg, cfg_negative)
else:
# --------------------------------
# kolors patch set conditioning
model, positive, negative, _ = patched_set_conds(model, positive, negative, None)
# --------------------------------
_guider, = self.get_custom_cls('BasicGuider').get_guider(model, positive)
# sampler
+7 -28
View File
@@ -1,21 +1,12 @@
import torch.nn
import comfy.model_management
import comfy.samplers
from comfy_extras.nodes_custom_sampler import Guider_Basic, Guider_DualCFG
_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):
from torch.nn import Linear
load_device = comfy.model_management.get_torch_device()
encoder_hid_proj_weight = sd.pop("encoder_hid_proj.weight")
encoder_hid_proj_bias = sd.pop("encoder_hid_proj.bias")
hid_proj = torch.nn.Linear(encoder_hid_proj_weight.shape[1], encoder_hid_proj_weight.shape[0])
hid_proj = Linear(encoder_hid_proj_weight.shape[1], encoder_hid_proj_weight.shape[0])
hid_proj.weight.data = encoder_hid_proj_weight
hid_proj.bias.data = encoder_hid_proj_bias
hid_proj = hid_proj.to(load_device)
@@ -24,6 +15,8 @@ def add_model_patch(model, sd):
}
def patched_set_conds(model, positive, negative=None, middle=None):
model_attrs = dir(model)
if "model_patch" in model.model_options:
mp = model.model_options["model_patch"]
if "hid_proj" in mp:
@@ -37,35 +30,21 @@ def patched_set_conds(model, positive, negative=None, middle=None):
positive[i][0] = hid_proj(positive[i][0])
if "control" in positive[i][1]:
if hasattr(positive[i][1]["control"], "control_model"):
positive[i][1]["control"].control_model.label_emb = model.model_patcher.model.diffusion_model.label_emb
positive[i][1]["control"].control_model.label_emb = model.model_patcher.model.diffusion_model.label_emb if "model_patcher" in model_attrs else model.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 = model.model_patcher.model.diffusion_model.label_emb
negative[i][1]["control"].control_model.label_emb = model.model_patcher.model.diffusion_model.label_emb if "model_patcher" in model_attrs else model.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 = model.model_patcher.model.diffusion_model.label_emb
middle[i][1]["control"].control_model.label_emb = model.model_patcher.model.diffusion_model.label_emb if "model_patcher" in model_attrs else model.model.diffusion_model.label_emb
return model, 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_cfgguider_set_conds
Guider_Basic.set_conds = patched_basicguider_set_conds
Guider_DualCFG.set_conds = patched_dualcfgguider_set_conds
+2 -1
View File
@@ -8,7 +8,6 @@ from comfy.model_patcher import ModelPatcher
from nodes import NODE_CLASS_MAPPINGS
from collections import defaultdict
from .log import log_node_info, log_node_error
from ..kolors.loader import load_chatglm3, applyKolorsUnet
from ..kolors.model_patch import add_model_patch
from ..dit.hunyuanDiT.loader import EXM_HyDiT_Tenc_Temp, load_hydit
from ..dit.pixArt.loader import load_pixart
@@ -485,6 +484,7 @@ class easyLoader:
log_node_info("Load Kolors UNet", f"{unet_name} cached")
return self.loaded_objects["unet"][unet_name][0]
else:
from ..kolors.loader import applyKolorsUnet
with applyKolorsUnet():
unet_path = folder_paths.get_full_path("unet", unet_name)
@@ -501,6 +501,7 @@ class easyLoader:
return model
def load_chatglm3(self, chatglm3_name):
from ..kolors.loader import load_chatglm3
if chatglm3_name in self.loaded_objects["chatglm3"]:
log_node_info("Load ChatGLM3", f"{chatglm3_name} cached")
return self.loaded_objects["chatglm3"][chatglm3_name][0]
+4 -6
View File
@@ -8,6 +8,7 @@ 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_set_conds
class easySampler:
def __init__(self):
@@ -108,9 +109,11 @@ class easySampler:
noise = comfy.sample.prepare_noise(latent_image, seed, batch_inds)
#######################################################################################
# add model patch
# brushnet
add_model_patch(model)
# model, positive, negative = patched_kolors_conds(model, positive, negative)
# kolors
model, positive, negative, _ = patched_set_conds(model, positive, negative, None)
#######################################################################################
samples = comfy.sample.sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative,
latent_image,
@@ -156,11 +159,6 @@ class easySampler:
preview_bytes = previewer.decode_latent_to_preview_image(preview_format, x0)
pbar.update_absolute(step + 1, total_steps, preview_bytes)
# samples = comfy.sample.sample_custom(model, noise, cfg, _sampler, sigmas, positive, negative, latent_image,
# noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar,
# seed=seed)
# 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)
+1 -1
View File
@@ -1,6 +1,6 @@
diffusers>=0.25.0
accelerate>=0.25.0
clip_interrogator>=0.6.0
diffusers>=0.25.0
lark-parser
onnxruntime
opencv-python