mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 01:04:32 +02:00
pulid with refine pass
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
@@ -19,20 +19,25 @@ cache = OrderedDict()
|
||||
embedder = None
|
||||
|
||||
|
||||
def prompt_compatible():
|
||||
def prompt_compatible(pipe = None):
|
||||
pipe = pipe or shared.sd_model
|
||||
if (
|
||||
'StableDiffusion' not in shared.sd_model.__class__.__name__ and
|
||||
'DemoFusion' not in shared.sd_model.__class__.__name__ and
|
||||
'StableCascade' not in shared.sd_model.__class__.__name__ and
|
||||
'Flux' not in shared.sd_model.__class__.__name__
|
||||
'StableDiffusion' not in pipe.__class__.__name__ and
|
||||
'DemoFusion' not in pipe.__class__.__name__ and
|
||||
'StableCascade' not in pipe.__class__.__name__ and
|
||||
'Flux' not in pipe.__class__.__name__
|
||||
):
|
||||
shared.log.warning(f"Prompt parser not supported: {shared.sd_model.__class__.__name__}")
|
||||
shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}")
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def prepare_model():
|
||||
pipe = shared.sd_model
|
||||
def prepare_model(pipe = None):
|
||||
pipe = pipe or shared.sd_model
|
||||
if not hasattr(pipe, "text_encoder") and hasattr(shared.sd_model, "pipe"):
|
||||
pipe = pipe.pipe
|
||||
if not hasattr(pipe, "text_encoder"):
|
||||
return None
|
||||
if shared.opts.diffusers_offload_mode == "balanced":
|
||||
pipe = sd_models.apply_balanced_offload(pipe)
|
||||
elif hasattr(pipe, "maybe_free_model_hooks"):
|
||||
@@ -62,7 +67,10 @@ class PromptEmbedder:
|
||||
earlyout = self.checkcache(p)
|
||||
if earlyout:
|
||||
return
|
||||
pipe = prepare_model()
|
||||
pipe = prepare_model(p.sd_model)
|
||||
if pipe is None:
|
||||
shared.log.error("Prompt encode: cannot find text encoder in model")
|
||||
return
|
||||
# per prompt in batch
|
||||
for batchidx, (prompt, negative_prompt) in enumerate(zip(self.prompts, self.negative_prompts)):
|
||||
self.prepare_schedule(prompt, negative_prompt)
|
||||
@@ -168,8 +176,8 @@ class PromptEmbedder:
|
||||
self.negative_pooleds[batchidx].append(negative_pooled)
|
||||
|
||||
if debug_enabled:
|
||||
get_tokens('positive', positive_prompt)
|
||||
get_tokens('negative', negative_prompt)
|
||||
get_tokens(pipe, 'positive', positive_prompt)
|
||||
get_tokens(pipe, 'negative', negative_prompt)
|
||||
pipe = prepare_model()
|
||||
|
||||
def __call__(self, key, step=0):
|
||||
@@ -288,25 +296,25 @@ def get_prompt_schedule(prompt, steps):
|
||||
return temp, len(schedule) > 1
|
||||
|
||||
|
||||
def get_tokens(msg, prompt):
|
||||
def get_tokens(pipe, msg, prompt):
|
||||
global token_dict, token_type # pylint: disable=global-statement
|
||||
if not shared.native:
|
||||
return 0
|
||||
if shared.sd_loaded and hasattr(shared.sd_model, 'tokenizer') and shared.sd_model.tokenizer is not None:
|
||||
if shared.sd_loaded and hasattr(pipe, 'tokenizer') and pipe.tokenizer is not None:
|
||||
if token_dict is None or token_type != shared.sd_model_type:
|
||||
token_type = shared.sd_model_type
|
||||
fn = shared.sd_model.tokenizer.name_or_path
|
||||
fn = pipe.tokenizer.name_or_path
|
||||
if fn.endswith('tokenizer'):
|
||||
fn = os.path.join(shared.sd_model.tokenizer.name_or_path, 'vocab.json')
|
||||
fn = os.path.join(pipe.tokenizer.name_or_path, 'vocab.json')
|
||||
else:
|
||||
fn = os.path.join(shared.sd_model.tokenizer.name_or_path, 'tokenizer', 'vocab.json')
|
||||
fn = os.path.join(pipe.tokenizer.name_or_path, 'tokenizer', 'vocab.json')
|
||||
token_dict = shared.readfile(fn, silent=True)
|
||||
for k, v in shared.sd_model.tokenizer.added_tokens_decoder.items():
|
||||
for k, v in pipe.tokenizer.added_tokens_decoder.items():
|
||||
token_dict[str(v)] = k
|
||||
shared.log.debug(f'Tokenizer: words={len(token_dict)} file="{fn}"')
|
||||
has_bos_token = shared.sd_model.tokenizer.bos_token_id is not None
|
||||
has_eos_token = shared.sd_model.tokenizer.eos_token_id is not None
|
||||
ids = shared.sd_model.tokenizer(prompt)
|
||||
has_bos_token = pipe.tokenizer.bos_token_id is not None
|
||||
has_eos_token = pipe.tokenizer.eos_token_id is not None
|
||||
ids = pipe.tokenizer(prompt)
|
||||
ids = getattr(ids, 'input_ids', [])
|
||||
tokens = []
|
||||
for i in ids:
|
||||
@@ -337,10 +345,10 @@ def normalize_prompt(pairs: list):
|
||||
return pairs
|
||||
|
||||
|
||||
def get_prompts_with_weights(prompt: str):
|
||||
def get_prompts_with_weights(pipe, prompt: str):
|
||||
t0 = time.time()
|
||||
manager = DiffusersTextualInversionManager(shared.sd_model, shared.sd_model.tokenizer or shared.sd_model.tokenizer_2)
|
||||
prompt = manager.maybe_convert_prompt(prompt, shared.sd_model.tokenizer or shared.sd_model.tokenizer_2)
|
||||
manager = DiffusersTextualInversionManager(pipe, pipe.tokenizer or pipe.tokenizer_2)
|
||||
prompt = manager.maybe_convert_prompt(prompt, pipe.tokenizer or pipe.tokenizer_2)
|
||||
texts_and_weights = prompt_parser.parse_prompt_attention(prompt)
|
||||
if shared.opts.prompt_mean_norm:
|
||||
texts_and_weights = normalize_prompt(texts_and_weights)
|
||||
@@ -348,7 +356,7 @@ def get_prompts_with_weights(prompt: str):
|
||||
if debug_enabled:
|
||||
all_tokens = 0
|
||||
for text in texts:
|
||||
tokens = get_tokens('section', text)
|
||||
tokens = get_tokens(pipe, 'section', text)
|
||||
all_tokens += tokens
|
||||
debug(f'Prompt tokenizer: parser={shared.opts.prompt_attention} tokens={all_tokens}')
|
||||
debug(f'Prompt: weights={texts_and_weights} time={(time.time() - t0):.3f}')
|
||||
@@ -412,7 +420,7 @@ def pad_to_same_length(pipe, embeds, empty_embedding_providers=None):
|
||||
return embeds
|
||||
|
||||
|
||||
def split_prompts(prompt, SD3 = False):
|
||||
def split_prompts(pipe, prompt, SD3 = False):
|
||||
if prompt.find("TE2:") != -1:
|
||||
prompt, prompt2 = prompt.split("TE2:")
|
||||
else:
|
||||
@@ -430,7 +438,7 @@ def split_prompts(prompt, SD3 = False):
|
||||
prompt3 = " " if prompt3.strip() == "" else prompt3.strip()
|
||||
|
||||
if SD3 and prompt3 != " ":
|
||||
ps, _ws = get_prompts_with_weights(prompt3)
|
||||
ps, _ws = get_prompts_with_weights(pipe, prompt3)
|
||||
prompt3 = " ".join(ps)
|
||||
return prompt, prompt2, prompt3
|
||||
|
||||
@@ -438,15 +446,15 @@ def split_prompts(prompt, SD3 = False):
|
||||
def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None):
|
||||
device = devices.device
|
||||
SD3 = hasattr(pipe, 'text_encoder_3')
|
||||
prompt, prompt_2, prompt_3 = split_prompts(prompt, SD3)
|
||||
neg_prompt, neg_prompt_2, neg_prompt_3 = split_prompts(neg_prompt, SD3)
|
||||
prompt, prompt_2, prompt_3 = split_prompts(pipe, prompt, SD3)
|
||||
neg_prompt, neg_prompt_2, neg_prompt_3 = split_prompts(pipe, neg_prompt, SD3)
|
||||
|
||||
if prompt != prompt_2:
|
||||
ps = [get_prompts_with_weights(p) for p in [prompt, prompt_2]]
|
||||
ns = [get_prompts_with_weights(p) for p in [neg_prompt, neg_prompt_2]]
|
||||
ps = [get_prompts_with_weights(pipe, p) for p in [prompt, prompt_2]]
|
||||
ns = [get_prompts_with_weights(pipe, p) for p in [neg_prompt, neg_prompt_2]]
|
||||
else:
|
||||
ps = 2 * [get_prompts_with_weights(prompt)]
|
||||
ns = 2 * [get_prompts_with_weights(neg_prompt)]
|
||||
ps = 2 * [get_prompts_with_weights(pipe, prompt)]
|
||||
ns = 2 * [get_prompts_with_weights(pipe, neg_prompt)]
|
||||
|
||||
positives, positive_weights = zip(*ps)
|
||||
negatives, negative_weights = zip(*ns)
|
||||
@@ -561,8 +569,8 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c
|
||||
|
||||
def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None):
|
||||
is_sd3 = hasattr(pipe, 'text_encoder_3')
|
||||
prompt, prompt_2, _prompt_3 = split_prompts(prompt, is_sd3)
|
||||
neg_prompt, neg_prompt_2, _neg_prompt_3 = split_prompts(neg_prompt, is_sd3)
|
||||
prompt, prompt_2, _prompt_3 = split_prompts(pipe, prompt, is_sd3)
|
||||
neg_prompt, neg_prompt_2, _neg_prompt_3 = split_prompts(pipe, neg_prompt, is_sd3)
|
||||
try:
|
||||
prompt = pipe.maybe_convert_prompt(prompt, pipe.tokenizer)
|
||||
neg_prompt = pipe.maybe_convert_prompt(neg_prompt, pipe.tokenizer)
|
||||
|
||||
Reference in New Issue
Block a user