mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
Merge pull request #2631 from AI-Casanova/prompt-callback
Prompt Scheduling for Diffusers
This commit is contained in:
@@ -231,6 +231,11 @@ 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 = []
|
||||
self.negative_pooleds = []
|
||||
|
||||
|
||||
@property
|
||||
|
||||
@@ -89,6 +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)
|
||||
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()
|
||||
@@ -293,36 +302,29 @@ 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
|
||||
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:
|
||||
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]
|
||||
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]
|
||||
else:
|
||||
args['negative_prompt'] = negative_prompts
|
||||
if hasattr(model, 'scheduler') and hasattr(model.scheduler, 'noise_sampler_seed') and hasattr(model.scheduler, 'noise_sampler'):
|
||||
@@ -342,7 +344,7 @@ def process_diffusers(p: StableDiffusionProcessing, seeds, prompts, negative_pro
|
||||
args['callback'] = diffusers_callback_legacy
|
||||
elif 'callback_on_step_end_tensor_inputs' in possible:
|
||||
args['callback_on_step_end'] = diffusers_callback
|
||||
args['callback_on_step_end_tensor_inputs'] = ['latents']
|
||||
args['callback_on_step_end_tensor_inputs'] = ['latents', 'prompt_embeds', 'negative_prompt_embeds']
|
||||
for arg in kwargs:
|
||||
if arg in possible: # add kwargs
|
||||
args[arg] = kwargs[arg]
|
||||
|
||||
@@ -58,32 +58,48 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager):
|
||||
debug(f'Prompt: expand={prompt}')
|
||||
return self.pipe.tokenizer.encode(prompt, add_special_tokens=False)
|
||||
|
||||
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, len(schedule) > 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, 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
|
||||
|
||||
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 = []
|
||||
|
||||
cache = {}
|
||||
for i in range(len(positive_schedule)):
|
||||
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:
|
||||
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):
|
||||
|
||||
Reference in New Issue
Block a user