pulid with refine pass

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2024-11-13 18:22:20 -05:00
parent c77370ef26
commit a0d55a5956
4 changed files with 53 additions and 35 deletions
+42 -34
View File
@@ -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)