support for sd3_medium_incl_clips in easy loader

This commit is contained in:
yolain
2024-06-13 02:41:05 +08:00
parent 5aa4f17187
commit b6b6bbfae4
5 changed files with 85 additions and 8 deletions
+1
View File
@@ -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**
+4 -3
View File
@@ -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)
+72 -1
View File
@@ -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
+5 -3
View File
@@ -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:
+3 -1
View File
@@ -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)