regional prompting

This commit is contained in:
Vladimir Mandic
2024-02-14 11:24:19 -05:00
parent d27295a923
commit ae4f7e6a5f
4 changed files with 42 additions and 15 deletions
+26 -3
View File
@@ -2,7 +2,21 @@
# https://github.com/huggingface/diffusers/blob/main/examples/community/regional_prompting_stable_diffusion.py
import gradio as gr
from modules import shared, devices, scripts, processing, sd_models
from diffusers.pipelines import pipeline_utils
from modules import shared, devices, scripts, processing, sd_models, prompt_parser_diffusers
def hijack_register_modules(self, **kwargs):
for name, module in kwargs.items():
if module is None or isinstance(module, (tuple, list)) and module[0] is None:
register_dict = {name: (None, None)}
elif isinstance(module, bool):
pass
else:
library, class_name = pipeline_utils._fetch_class_library_tuple(module) # pylint: disable=protected-access
register_dict = {name: (library, class_name)}
self.register_to_config(**register_dict)
setattr(self, name, module)
class Script(scripts.Script):
@@ -10,7 +24,6 @@ class Script(scripts.Script):
return 'Regional prompting'
def show(self, is_img2img):
return False
return not is_img2img if shared.backend == shared.Backend.DIFFUSERS else False
def change(self, mode):
@@ -39,6 +52,10 @@ class Script(scripts.Script):
if shared.sd_model_type != 'sd':
shared.log.error(f'Regional prompting: incorrect base model: {shared.sd_model.__class__.__name__}')
return
pipeline_utils.DiffusionPipeline.register_modules = hijack_register_modules
prompt_parser_diffusers.EmbeddingsProvider._encode_token_ids_to_embeddings = prompt_parser_diffusers.orig_encode_token_ids_to_embeddings # pylint: disable=protected-access
shared.sd_model = sd_models.switch_pipe('regional_prompting_stable_diffusion', shared.sd_model)
if shared.sd_model.__class__.__name__ != 'RegionalPromptingStableDiffusionPipeline': # switch failed
shared.log.error(f'Regional prompting: not a tiling pipeline: {shared.sd_model.__class__.__name__}')
@@ -56,11 +73,17 @@ class Script(scripts.Script):
rp_args['th'] = threshold
else:
rp_args['div'] = grid
p.task_args = { **p.task_args, 'rp_args': rp_args }
p.task_args = {
**p.task_args,
'prompt': p.prompt,
'rp_args': rp_args,
}
# run pipeline
shared.log.debug(f'Regional: args={p.task_args}')
processed: processing.Processed = processing.process_images(p) # runs processing using main loop
# restore pipeline and params
prompt_parser_diffusers.EmbeddingsProvider._encode_token_ids_to_embeddings = prompt_parser_diffusers.compel_hijack # pylint: disable=protected-access
shared.opts.data['prompt_attention'] = orig_prompt_attention
shared.sd_model = orig_pipeline
shared.sd_model.to(orig_dtype)