mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
regional prompting
This commit is contained in:
@@ -26,6 +26,10 @@
|
||||
- **Clip-skip** reworked completely, thanks @AI-Casanova & @Disty0
|
||||
now clip-skip range is 0-12 where previously lowest value was 1 (default is still 1)
|
||||
values can also be decimal to interpolate between different layers, for example `clip-skip: 1.5`, thanks @AI-Casanova
|
||||
- **Regional prompting** as a built-in solution
|
||||
simply enable from scripts -> regional prompting
|
||||
note that usage is same as original implementation from @hako-mikan
|
||||
click on title to open docs and see examples of full syntax on how to use it
|
||||
- **Cross-attention** refactored cross-attention methods, thanks @Disty0
|
||||
- for backend:original, its unchanged: SDP, xFormers, Doggettxs, InvokeAI, Sub-quadratic, Split attention
|
||||
- for backend:diffuers, list is now: SDP, xFormers, Batch matrix-matrix, Split attention, Dynamic Attention BMM, Dynamic Attention SDP
|
||||
|
||||
@@ -426,8 +426,6 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
|
||||
p.extra_generation_params["Sampler Eta"] = shared.opts.scheduler_eta
|
||||
try:
|
||||
t0 = time.time()
|
||||
img = base_args['mask_image']
|
||||
img.save('/tmp/mask.png')
|
||||
output = shared.sd_model(**base_args) # pylint: disable=not-callable
|
||||
if isinstance(output, dict):
|
||||
output = SimpleNamespace(**output)
|
||||
|
||||
@@ -7,22 +7,26 @@ from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsPr
|
||||
from transformers import PreTrainedTokenizer
|
||||
from modules import shared, prompt_parser, devices
|
||||
|
||||
|
||||
debug = shared.log.trace if os.environ.get('SD_PROMPT_DEBUG', None) is not None else lambda *args, **kwargs: None
|
||||
debug('Trace: PROMPT')
|
||||
orig_encode_token_ids_to_embeddings = EmbeddingsProvider._encode_token_ids_to_embeddings # pylint: disable=protected-access
|
||||
|
||||
|
||||
def compel_hijack(self, token_ids: torch.Tensor,
|
||||
attention_mask: typing.Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
needs_hidden_states = self.returned_embeddings_type != 1
|
||||
text_encoder_output = self.text_encoder(token_ids,
|
||||
attention_mask,
|
||||
output_hidden_states=needs_hidden_states,
|
||||
return_dict=True)
|
||||
text_encoder_output = self.text_encoder(token_ids, attention_mask, output_hidden_states=needs_hidden_states, return_dict=True)
|
||||
if not needs_hidden_states:
|
||||
return text_encoder_output.last_hidden_state
|
||||
normalized = self.returned_embeddings_type > 0
|
||||
clip_skip = math.floor(abs(self.returned_embeddings_type))
|
||||
interpolation = abs(self.returned_embeddings_type) - clip_skip
|
||||
try:
|
||||
normalized = self.returned_embeddings_type > 0
|
||||
clip_skip = math.floor(abs(self.returned_embeddings_type))
|
||||
interpolation = abs(self.returned_embeddings_type) - clip_skip
|
||||
except Exception:
|
||||
normalized = False
|
||||
clip_skip = 1
|
||||
interpolation = False
|
||||
if interpolation:
|
||||
hidden_state = (1 - interpolation) * text_encoder_output.hidden_states[-clip_skip] + interpolation * text_encoder_output.hidden_states[-(clip_skip+1)]
|
||||
else:
|
||||
@@ -32,9 +36,7 @@ def compel_hijack(self, token_ids: torch.Tensor,
|
||||
return hidden_state
|
||||
|
||||
|
||||
|
||||
|
||||
EmbeddingsProvider._encode_token_ids_to_embeddings = compel_hijack
|
||||
EmbeddingsProvider._encode_token_ids_to_embeddings = compel_hijack # pylint: disable=protected-access
|
||||
|
||||
|
||||
# from https://github.com/damian0815/compel/blob/main/src/compel/diffusers_textual_inversion_manager.py
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user