From b6b6bbfae4d56e7a12f7ec765fdf9da9de20a69c Mon Sep 17 00:00:00 2001 From: yolain Date: Thu, 13 Jun 2024 02:41:05 +0800 Subject: [PATCH] support for sd3_medium_incl_clips in easy loader --- README.md | 1 + py/easyNodes.py | 7 +++-- py/libs/adv_encode.py | 73 ++++++++++++++++++++++++++++++++++++++++++- py/libs/sampler.py | 8 +++-- py/libs/utils.py | 4 ++- 5 files changed, 85 insertions(+), 8 deletions(-) diff --git a/README.md b/README.md index 692048b..c594138 100644 --- a/README.md +++ b/README.md @@ -51,6 +51,7 @@ git clone https://github.com/yolain/ComfyUI-Easy-Use **v1.1.9** +- 支持 sd3_medium_incl_clips 模型 - 增加 `easy applyInpaint` - 局部重绘全模式节点 (相比与之前的kSamplerInpating节点逻辑会更合理些) **v1.1.8** diff --git a/py/easyNodes.py b/py/easyNodes.py index d8ad41a..bd0400e 100644 --- a/py/easyNodes.py +++ b/py/easyNodes.py @@ -938,13 +938,14 @@ class fullLoader: # Clean models from loaded_objects easyCache.update_loaded_objects(prompt) - # Create Empty Latent - samples = sampler.emptyLatent(resolution, empty_latent_width, empty_latent_height, batch_size) - # Load models log_node_warn("正在加载模型...") model, clip, vae, clip_vision, lora_stack = easyCache.load_main(ckpt_name, config_name, vae_name, lora_name, lora_model_strength, lora_clip_strength, optional_lora_stack, model_override, clip_override, vae_override, prompt) + # Create Empty Latent + sd3 = True if get_sd_version(model) == 'sd3' else False + samples = sampler.emptyLatent(resolution, empty_latent_width, empty_latent_height, batch_size, sd3=sd3) + # Prompt to Conditioning positive_embeddings_final, positive_wildcard_prompt, model, clip = prompt_to_cond('positive', model, clip, clip_skip, lora_stack, positive, positive_token_normalization, positive_weight_interpretation, a1111_prompt_style, my_unique_id, prompt, easyCache) negative_embeddings_final, negative_wildcard_prompt, model, clip = prompt_to_cond('negative', model, clip, clip_skip, lora_stack, negative, negative_token_normalization, negative_weight_interpretation, a1111_prompt_style, my_unique_id, prompt, easyCache) diff --git a/py/libs/adv_encode.py b/py/libs/adv_encode.py index fbebf48..b2b2846 100644 --- a/py/libs/adv_encode.py +++ b/py/libs/adv_encode.py @@ -4,6 +4,11 @@ import itertools from comfy import model_management from comfy.sdxl_clip import SDXLClipModel, SDXLRefinerClipModel, SDXLClipG +try: + from comfy.sd3_clip import SD3ClipModel, SDT5XXLModel +except: + SD3ClipModel, SDT5XXLModel = None, None + pass from nodes import NODE_CLASS_MAPPINGS, ConditioningConcat def _grouper(n, iterable): @@ -238,6 +243,9 @@ def encode_token_weights_l(model, token_weight_pairs): l_out, pooled = model.clip_l.encode_token_weights(token_weight_pairs) return l_out, pooled +def encode_token_weights_t5(model, token_weight_pairs): + return model.t5xxl.encode_token_weights(token_weight_pairs) + def encode_token_weights(model, token_weight_pairs, encode_func): if model.layer_idx is not None: @@ -258,6 +266,14 @@ def prepareXL(embs_l, embs_g, pooled, clip_balance): else: return embs_g, pooled +def prepareSD3(out, pooled, clip_balance): + lg_w = 1 - max(0, clip_balance - .5) * 2 + t5_w = 1 - max(0, .5 - clip_balance) * 2 + if out.shape[0] > 1: + return torch.cat([out[0] * lg_w, out[1] * t5_w], dim=-1), pooled + else: + return out, pooled + def advanced_encode(clip, text, token_normalization, weight_interpretation, w_max=1.0, clip_balance=.5, apply_to_pooled=True, width=1024, height=1024, crop_w=0, crop_h=0, target_width=1024, target_height=1024, a1111_prompt_style=False, steps=1): @@ -283,7 +299,62 @@ def advanced_encode(clip, text, token_normalization, weight_interpretation, w_ma for text in pass3: tokenized = clip.tokenize(text, return_word_ids=True) - if isinstance(clip.cond_stage_model, (SDXLClipModel, SDXLRefinerClipModel, SDXLClipG)): + if SD3ClipModel and isinstance(clip.cond_stage_model, SD3ClipModel): + lg_out = None + pooled = None + out = None + + if len(tokenized['l']) > 0 or len(tokenized['g']) > 0: + if 'l' in tokenized: + lg_out, l_pooled = advanced_encode_from_tokens(tokenized['l'], + token_normalization, + weight_interpretation, + lambda x: encode_token_weights(clip, x, encode_token_weights_l), + w_max=w_max, return_pooled=True,) + else: + l_pooled = torch.zeros((1, 768), device=model_management.intermediate_device()) + + if 'g' in tokenized: + g_out, g_pooled = advanced_encode_from_tokens(tokenized['g'], + token_normalization, + weight_interpretation, + lambda x: encode_token_weights(clip, x, encode_token_weights_g), + w_max=w_max, return_pooled=True) + if lg_out is not None: + lg_out = torch.cat([lg_out, g_out], dim=-1) + else: + lg_out = torch.nn.functional.pad(g_out, (768, 0)) + else: + g_out = None + g_pooled = torch.zeros((1, 1280), device=model_management.intermediate_device()) + + if lg_out is not None: + lg_out = torch.nn.functional.pad(lg_out, (0, 4096 - lg_out.shape[-1])) + out = lg_out + pooled = torch.cat((l_pooled, g_pooled), dim=-1) + + # t5xxl not working, need to fix + if 't5xxl' in tokenized and clip.cond_stage_model.t5xxl is not None: + t5_out, t5_pooled = advanced_encode_from_tokens(tokenized['t5xxl'], + token_normalization, + weight_interpretation, + lambda x: encode_token_weights(clip, x, encode_token_weights_t5), + w_max=w_max, return_pooled=True) + if lg_out is not None: + out = torch.cat([lg_out, t5_out], dim=-2) + else: + out = t5_out + + if out is None: + out = torch.zeros((1, 77, 4096), device=model_management.intermediate_device()) + + if pooled is None: + pooled = torch.zeros((1, 768 + 1280), device=model_management.intermediate_device()) + + embeddings_final, pooled = prepareSD3(out, pooled, clip_balance) + cond = [[embeddings_final, {"pooled_output": pooled}]] + + elif isinstance(clip.cond_stage_model, (SDXLClipModel, SDXLRefinerClipModel, SDXLClipG)): embs_l = None embs_g = None pooled = None diff --git a/py/libs/sampler.py b/py/libs/sampler.py index 107e937..b7aee37 100644 --- a/py/libs/sampler.py +++ b/py/libs/sampler.py @@ -51,7 +51,7 @@ class easySampler: parts.append('None') return parts - def emptyLatent(self, resolution, empty_latent_width, empty_latent_height, batch_size=1, compression=0): + def emptyLatent(self, resolution, empty_latent_width, empty_latent_height, batch_size=1, compression=0, sd3=False): if resolution != "自定义 x 自定义": try: width, height = map(int, resolution.split(' x ')) @@ -59,8 +59,10 @@ class easySampler: empty_latent_height = height except ValueError: raise ValueError("Invalid base_resolution format.") - - if compression == 0: + if sd3: + latent = torch.ones([batch_size, 16, empty_latent_height // 8, empty_latent_width // 8], device=self.device) * 0.0609 + samples = {"samples": latent} + elif compression == 0: latent = torch.zeros([batch_size, 4, empty_latent_height // 8, empty_latent_width // 8], device=self.device) samples = {"samples": latent} else: diff --git a/py/libs/utils.py b/py/libs/utils.py index f53fe70..62d1a1e 100644 --- a/py/libs/utils.py +++ b/py/libs/utils.py @@ -90,7 +90,9 @@ import comfy.supported_models_base def get_sd_version(model): base: BaseModel = model.model model_config: comfy.supported_models.supported_models_base.BASE = base.model_config - if isinstance(model_config, comfy.supported_models.SDXL): + if isinstance(model_config, comfy.supported_models.SD3): + return 'sd3' + elif isinstance(model_config, comfy.supported_models.SDXL): return 'sdxl' elif isinstance( model_config, (comfy.supported_models.SD15, comfy.supported_models.SD20)