diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 1fe11b817..417139bde 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -89,6 +89,8 @@ def process_diffusers(p: StableDiffusionProcessing, seeds, prompts, negative_pro if kwargs.get('latents', None) is None: return kwargs kwargs = correction_callback(p, timestep, kwargs) + kwargs["prompt_embeds"] = p.prompt_embeds[step-1] + kwargs["negative_prompt_embeds"] = p.negative_embeds[step-1] shared.state.current_latent = kwargs['latents'] if shared.cmd_opts.profile and shared.profiler is not None: shared.profiler.step() @@ -293,36 +295,37 @@ def process_diffusers(p: StableDiffusionProcessing, seeds, prompts, negative_pro possible = signature.parameters.keys() generator_device = devices.cpu if shared.opts.diffusers_generator_device == "cpu" else shared.device generator = [torch.Generator(generator_device).manual_seed(s) for s in seeds] - prompt_embed = None - pooled = None - negative_embed = None - negative_pooled = None + # prompt_embed = None + # pooled = None + # negative_embed = None + # negative_pooled = None prompts, negative_prompts, prompts_2, negative_prompts_2 = fix_prompts(prompts, negative_prompts, prompts_2, negative_prompts_2) parser = 'Fixed attention' if shared.opts.prompt_attention != 'Fixed attention' and 'StableDiffusion' in model.__class__.__name__: try: - prompt_embed, pooled, negative_embed, negative_pooled = prompt_parser_diffusers.encode_prompts(model, prompts, negative_prompts, kwargs.pop("clip_skip", None)) + prompt_parser_diffusers.encode_prompts(model, p, prompts, negative_prompts, kwargs.get("num_inference_steps", 1), 0, kwargs.pop("clip_skip", None)) + # prompt_embed, pooled, negative_embed, negative_pooled = , , , , parser = shared.opts.prompt_attention except Exception as e: shared.log.error(f'Prompt parser encode: {e}') if os.environ.get('SD_PROMPT_DEBUG', None) is not None: errors.display(e, 'Prompt parser encode') if 'prompt' in possible: - if hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and prompt_embed is not None: - if type(pooled) == list: - pooled = pooled[0] - if type(negative_pooled) == list: - negative_pooled = negative_pooled[0] - args['prompt_embeds'] = prompt_embed + if hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and p.prompt_embeds[0] is not None: + # if type(pooled) == list: + # pooled = pooled[0] + # if type(negative_pooled) == list: + # negative_pooled = p.negative_pooleds[0][0] + args['prompt_embeds'] = p.prompt_embeds[0] if 'XL' in model.__class__.__name__: - args['pooled_prompt_embeds'] = pooled + args['pooled_prompt_embeds'] = p.positive_pooleds[0][0] else: args['prompt'] = prompts if 'negative_prompt' in possible: - if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and negative_embed is not None: - args['negative_prompt_embeds'] = negative_embed + if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and p.negative_embeds[0] is not None: + args['negative_prompt_embeds'] = p.negative_embeds[0] if 'XL' in model.__class__.__name__: - args['negative_pooled_prompt_embeds'] = negative_pooled + args['negative_pooled_prompt_embeds'] = p.negative_pooleds[0][0] else: args['negative_prompt'] = negative_prompts if hasattr(model, 'scheduler') and hasattr(model.scheduler, 'noise_sampler_seed') and hasattr(model.scheduler, 'noise_sampler'): diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index f69b7e71e..ceba53977 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -58,32 +58,43 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager): debug(f'Prompt: expand={prompt}') return self.pipe.tokenizer.encode(prompt, add_special_tokens=False) +def get_prompt_schedule(prompt, steps, step): + temp = [] + schedule = prompt_parser.get_learned_conditioning_prompt_schedules([prompt], steps)[0] + for chunk in schedule: + for s in range(steps): + if s + 1 <= chunk[0]: + temp.append(chunk[1]) + return temp[step-1] -def encode_prompts(pipe, prompts: list, negative_prompts: list, clip_skip: typing.Optional[int] = None): +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: - prompt_embeds = [] - positive_pooleds = [] - negative_embeds = [] - negative_pooleds = [] - for i in range(len(prompts)): - prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, prompts[i], negative_prompts[i], clip_skip) - prompt_embeds.append(prompt_embed) - positive_pooleds.append(positive_pooled) - negative_embeds.append(negative_embed) - negative_pooleds.append(negative_pooled) + positive_schedule = get_prompt_schedule(prompts[0]) + negative_schedule = get_prompt_schedule(negative_prompts[0]) - if prompt_embeds is not None: - prompt_embeds = torch.cat(prompt_embeds, dim=0) - if negative_embeds is not None: - negative_embeds = torch.cat(negative_embeds, dim=0) - if positive_pooleds is not None and shared.sd_model_type == "sdxl": - positive_pooleds = torch.cat(positive_pooleds, dim=0) - if negative_pooleds is not None and shared.sd_model_type == "sdxl": - negative_pooleds = torch.cat(negative_pooleds, dim=0) - return prompt_embeds, positive_pooleds, negative_embeds, negative_pooleds + p.prompt_embeds = [] + p.positive_pooleds = [] + p.negative_embeds = [] + p.negative_pooleds = [] + + 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) + + if prompt_embed is not None: + p.prompt_embeds.append(torch.cat([prompt_embed]*len(prompts), dim=0)) + if negative_embed is not None: + p.negative_embeds.append(torch.cat([negative_embed]*len(negative_prompts), dim=0)) + if positive_pooled is not None and shared.sd_model_type == "sdxl": + p.positive_pooleds.append(torch.cat([positive_pooled]*len(prompts), dim=0)) + if negative_pooled is not None and shared.sd_model_type == "sdxl": + p.negative_pooleds.append(torch.cat([negative_pooled]*len(negative_prompts), dim=0)) + return def get_prompts_with_weights(prompt: str):