Move embedder object, cleanup stepwise lora

This commit is contained in:
AI-Casanova
2024-10-30 23:24:50 -05:00
parent 38303f0c61
commit 3932d5fd1b
4 changed files with 27 additions and 26 deletions
+1
View File
@@ -104,6 +104,7 @@ def activate(p, extra_network_data, step=0):
p.extra_network_data = extra_network_data
if stepwise:
p.stepwise_lora = True
shared.opts.data['lora_functional'] = functional
+18 -18
View File
@@ -117,7 +117,7 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2
'Flux' in model.__class__.__name__
):
try:
p.embedder = prompt_parser_diffusers.PromptEmbedder(prompts, negative_prompts, clip_skip, p)
prompt_parser_diffusers.embedder = prompt_parser_diffusers.PromptEmbedder(prompts, negative_prompts, steps, clip_skip, p)
parser = shared.opts.prompt_attention
except Exception as e:
shared.log.error(f'Prompt parser encode: {e}')
@@ -128,27 +128,27 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2
if 'prompt' in possible:
if 'OmniGen' in model.__class__.__name__:
prompts = [p.replace('|image|', '<|image_1|>') for p in prompts]
if hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and p.embedder is not None:
args['prompt_embeds'] = p.embedder('prompt_embeds')
if hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
args['prompt_embeds'] = prompt_parser_diffusers.embedder('prompt_embeds')
if 'StableCascade' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0:
args['prompt_embeds_pooled'] = p.embedder('positive_pooleds').unsqueeze(0)
elif 'XL' in model.__class__.__name__ and p.embedder is not None:
args['pooled_prompt_embeds'] = p.embedder('positive_pooleds')
elif 'StableDiffusion3' in model.__class__.__name__ and p.embedder is not None:
args['pooled_prompt_embeds'] = p.embedder('positive_pooleds')
elif 'Flux' in model.__class__.__name__ and p.embedder is not None:
args['pooled_prompt_embeds'] = p.embedder('positive_pooleds')
args['prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('positive_pooleds').unsqueeze(0)
elif 'XL' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
elif 'StableDiffusion3' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
elif 'Flux' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
else:
args['prompt'] = prompts
if 'negative_prompt' in possible:
if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and p.embedder is not None:
args['negative_prompt_embeds'] = p.embedder('negative_prompt_embeds')
if 'StableCascade' in model.__class__.__name__ and p.embedder is not None:
args['negative_prompt_embeds_pooled'] = p.embedder('negative_pooleds').unsqueeze(0)
if 'XL' in model.__class__.__name__ and p.embedder is not None:
args['negative_pooled_prompt_embeds'] = p.embedder('negative_pooleds')
if 'StableDiffusion3' in model.__class__.__name__ and p.embedder is not None:
args['negative_pooled_prompt_embeds'] = p.embedder('negative_pooleds')
if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
args['negative_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_prompt_embeds')
if 'StableCascade' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
args['negative_prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('negative_pooleds').unsqueeze(0)
if 'XL' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
if 'StableDiffusion3' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
else:
if 'PixArtSigmaPipeline' in model.__class__.__name__: # pixart-sigma pipeline throws list-of-list for negative prompt
args['negative_prompt'] = negative_prompts[0]
+5 -6
View File
@@ -3,8 +3,7 @@ import os
import time
import torch
import numpy as np
from modules import shared, processing_correction, extra_networks, timer
from modules import shared, processing_correction, extra_networks, timer, prompt_parser_diffusers
p = None
debug_callback = shared.log.trace if os.environ.get('SD_CALLBACK_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -49,7 +48,7 @@ def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict):
if shared.state.interrupted or shared.state.skipped:
raise AssertionError('Interrupted...')
time.sleep(0.1)
if hasattr(p, "extra_network_data"):
if hasattr(p, "stepwise_lora"):
extra_networks.activate(p, p.extra_network_data, step=step)
if latents is None:
return kwargs
@@ -67,12 +66,12 @@ def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict):
pipe.set_ip_adapter_scale(ip_adapter_scales)
if step != getattr(pipe, 'num_timesteps', 0):
kwargs = processing_correction.correction_callback(p, timestep, kwargs)
if p.embedder is not None:
if prompt_parser_diffusers.embedder is not None:
try:
if 'prompt_embeds' in kwargs:
kwargs["prompt_embeds"] = p.embedder("prompt_embeds", step + 1)
kwargs["prompt_embeds"] = prompt_parser_diffusers.embedder("prompt_embeds", step + 1)
if 'negative_prompt_embeds' in kwargs:
kwargs["negative_prompt_embeds"] = p.embedder("negative_prompt_embeds", step + 1)
kwargs["negative_prompt_embeds"] = prompt_parser_diffusers.embedder("negative_prompt_embeds", step + 1)
except Exception as e:
shared.log.debug(f"Callback: {e}")
if step == int(getattr(pipe, 'num_timesteps', 100) * p.cfg_end) and 'prompt_embeds' in kwargs and 'negative_prompt_embeds' in kwargs:
+3 -2
View File
@@ -16,6 +16,7 @@ orig_encode_token_ids_to_embeddings = EmbeddingsProvider._encode_token_ids_to_em
token_dict = None # used by helper get_tokens
token_type = None # used by helper get_tokens
cache = OrderedDict()
embedder = None
def prompt_compatible():
@@ -41,13 +42,13 @@ def prepare_model():
class PromptEmbedder:
def __init__(self, prompts, negative_prompts, clip_skip, p):
def __init__(self, prompts, negative_prompts, steps, clip_skip, p):
t0 = time.time()
self.prompts = prompts
self.negative_prompts = negative_prompts
self.batchsize = len(self.prompts)
self.allsame = self.compare_prompts() # collapses batched prompts to single prompt if possible
self.steps = p.steps
self.steps = steps
self.clip_skip = clip_skip
# All embeds are nested lists, outer list batch length, inner schedule length
self.prompt_embeds = [[]] * self.batchsize