diff --git a/modules/processing.py b/modules/processing.py index 68a1912c3..6d51d6d1d 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -231,6 +231,7 @@ class StableDiffusionProcessing: self.hdr_maximize = hdr_maximize self.hdr_max_center = hdr_max_center self.hdr_max_boundry = hdr_max_boundry + self.scheduled_prompt: bool = False self.prompt_embeds = [] self.positive_pooleds = [] self.negative_embeds = [] diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 25f6fe646..66c607fe9 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -89,13 +89,15 @@ def process_diffusers(p: StableDiffusionProcessing, seeds, prompts, negative_pro if kwargs.get('latents', None) is None: return kwargs kwargs = correction_callback(p, timestep, kwargs) - try: - kwargs["prompt_embeds"] = p.prompt_embeds[step + 1].repeat(1, kwargs["prompt_embeds"].shape[0], 1).view( - kwargs["prompt_embeds"].shape[0], kwargs["prompt_embeds"].shape[1], -1) - kwargs["negative_prompt_embeds"] = p.negative_embeds[step + 1].repeat(1, kwargs["negative_prompt_embeds"].shape[0], 1).view( - kwargs["negative_prompt_embeds"].shape[0], kwargs["negative_prompt_embeds"].shape[1], -1) - except: - pass + if p.scheduled_prompt: + try: + i = (step + 1) % len(p.prompt_embeds) + kwargs["prompt_embeds"] = p.prompt_embeds[i][0:1].repeat(1, kwargs["prompt_embeds"].shape[0], 1).view( + kwargs["prompt_embeds"].shape[0], kwargs["prompt_embeds"].shape[1], -1) + kwargs["negative_prompt_embeds"] = p.negative_embeds[i][0:1].repeat(1, kwargs["negative_prompt_embeds"].shape[0], 1).view( + kwargs["negative_prompt_embeds"].shape[0], kwargs["negative_prompt_embeds"].shape[1], -1) + except Exception as e: + shared.log.debug(f"Callback: {e}") shared.state.current_latent = kwargs['latents'] if shared.cmd_opts.profile and shared.profiler is not None: shared.profiler.step() diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 16fab82d3..ae5acb8f6 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -58,33 +58,39 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager): debug(f'Prompt: expand={prompt}') return self.pipe.tokenizer.encode(prompt, add_special_tokens=False) -def get_prompt_schedule(prompt, steps): +def get_prompt_schedule(p, prompt, steps): temp = [] schedule = prompt_parser.get_learned_conditioning_prompt_schedules([prompt], steps)[0] for chunk in schedule: for s in range(steps): if len(temp) < s + 1 <= chunk[0]: temp.append(chunk[1]) - return temp + return temp, len(schedule) > 1 def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, step: int = 1, clip_skip: typing.Optional[int] = None): if 'StableDiffusion' not in pipe.__class__.__name__ and 'DemoFusion': shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}") return None, None, None, None else: - positive_schedule = get_prompt_schedule(prompts[0], steps) - negative_schedule = get_prompt_schedule(negative_prompts[0], steps) + positive_schedule, scheduled = get_prompt_schedule(p, prompts[0], steps) + negative_schedule, neg_scheduled = get_prompt_schedule(p, negative_prompts[0], steps) + p.scheduled_prompt = scheduled or neg_scheduled p.prompt_embeds = [] p.positive_pooleds = [] p.negative_embeds = [] p.negative_pooleds = [] + cache = {} for i in range(len(positive_schedule)): - prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, - positive_schedule[i], - negative_schedule[i], - clip_skip) + cached = cache.get(positive_schedule[i]+negative_schedule[i], None) + if cached is not None: + prompt_embed, positive_pooled, negative_embed, negative_pooled = cached + else: + prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, + positive_schedule[i], + negative_schedule[i], + clip_skip) if prompt_embed is not None: p.prompt_embeds.append(torch.cat([prompt_embed]*len(prompts), dim=0)) if negative_embed is not None: