support for sd3_medium_incl_clips in easy loader
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user