From 3932d5fd1b2f69a119217871bf6b286321fdcac5 Mon Sep 17 00:00:00 2001 From: AI-Casanova <54461896+AI-Casanova@users.noreply.github.com> Date: Wed, 30 Oct 2024 23:24:50 -0500 Subject: [PATCH] Move embedder object, cleanup stepwise lora --- modules/extra_networks.py | 1 + modules/processing_args.py | 36 +++++++++++++++--------------- modules/processing_callbacks.py | 11 +++++---- modules/prompt_parser_diffusers.py | 5 +++-- 4 files changed, 27 insertions(+), 26 deletions(-) diff --git a/modules/extra_networks.py b/modules/extra_networks.py index 673549b6b..b464bd349 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -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 diff --git a/modules/processing_args.py b/modules/processing_args.py index 34dd97a1c..5cdf290a7 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -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] diff --git a/modules/processing_callbacks.py b/modules/processing_callbacks.py index 5c24aead0..3ace64ed8 100644 --- a/modules/processing_callbacks.py +++ b/modules/processing_callbacks.py @@ -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: diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 5ba0e8a74..678649a66 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -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