From c68258304c88708d9ef60e364dec6257c013dc2c Mon Sep 17 00:00:00 2001 From: yolain Date: Sat, 13 Jul 2024 22:48:43 +0800 Subject: [PATCH] Fix some kolors logic --- py/easyNodes.py | 15 ++++++++++++++- py/kolors/model_patch.py | 35 +++++++---------------------------- py/libs/loader.py | 3 ++- py/libs/sampler.py | 10 ++++------ requirements.txt | 2 +- 5 files changed, 28 insertions(+), 37 deletions(-) diff --git a/py/easyNodes.py b/py/easyNodes.py index 87ceb4d..e2ba222 100644 --- a/py/easyNodes.py +++ b/py/easyNodes.py @@ -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 diff --git a/py/kolors/model_patch.py b/py/kolors/model_patch.py index 866d0af..63f1a2f 100644 --- a/py/kolors/model_patch.py +++ b/py/kolors/model_patch.py @@ -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 - diff --git a/py/libs/loader.py b/py/libs/loader.py index fe4c4ac..46a3c28 100644 --- a/py/libs/loader.py +++ b/py/libs/loader.py @@ -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] diff --git a/py/libs/sampler.py b/py/libs/sampler.py index 3163955..e5982a6 100644 --- a/py/libs/sampler.py +++ b/py/libs/sampler.py @@ -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) diff --git a/requirements.txt b/requirements.txt index 341c6cb..cc1920c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -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