Fix some kolors logic
This commit is contained in:
+14
-1
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user