diff --git a/TODO.md b/TODO.md index 986808f4f..d50ece4db 100644 --- a/TODO.md +++ b/TODO.md @@ -2,15 +2,14 @@ ## Release -- Review: feat/tag-dictionaries - Test: Zeta-Chroma - Test: Anima-Preview-3 - Test: SDXS-1B - Test: VIBE-Image-Edit -- Test: Bria-FIBO - Test: Lumina-DiMOO - Test: Step1X-Edit -- Code: prompt encode for Bria-FIBO: +- Test: Bria-FIBO prompt to json +- Test: Bria-FIBO - Port: ERNIE-Image (merged, unpublished) - Port: NucleusMoE-Image (merged, unpublished) - Port: JoyAI-Image-Edit (in-progress, published) diff --git a/modules/sd_hijack_te.py b/modules/sd_hijack_te.py index b35746cfa..0edf11d24 100644 --- a/modules/sd_hijack_te.py +++ b/modules/sd_hijack_te.py @@ -13,7 +13,14 @@ def hijack_encode_prompt(*args, **kwargs): prompt = kwargs.get('prompt', None) or (args[0] if len(args) > 0 else None) if prompt is not None: log.debug(f'Encode: prompt="{prompt}" hijack=True') - res = shared.sd_model.orig_encode_prompt(*args, **kwargs) + if hasattr(shared.sd_model, 'before_prompt_encode'): + prompt = shared.sd_model.before_prompt_encode(prompt) + if hasattr(shared.sd_model, 'orig_encode_prompt'): + res = shared.sd_model.orig_encode_prompt(*args, **kwargs) + if hasattr(shared.sd_model, 'after_prompt_encode'): + res = shared.sd_model.after_prompt_encode(res) + else: + res = prompt except Exception as e: log.error(f'Encode prompt: {e}') errors.display(e, 'Encode prompt') diff --git a/pipelines/bria/prompt_to_json.py b/pipelines/bria/prompt_to_json.py new file mode 100644 index 000000000..91a944b2e --- /dev/null +++ b/pipelines/bria/prompt_to_json.py @@ -0,0 +1,48 @@ +import time +import json + + +model = None + + +def generate_prompt_local(prompt, image=None, repo_id="briaai/FIBO-edit-prompt-to-JSON"): + from modules import shared, devices, sd_models, model_quant + from modules.logger import log + + global model # pylint: disable=global-statement + if model is None: + from diffusers.modular_pipelines import ModularPipelineBlocks + quant_args = model_quant.create_config(module='LLM') + pipeline = ModularPipelineBlocks.from_pretrained(repo_id, + trust_remote_code=True, + torch_dtype=devices.dtype, + cache_dir=shared.opts.hfcache_dir, + **quant_args, + ) + model = pipeline.init_pipeline() + log.debug(f'JSONEncode loaded: model={model} cls={model.__class__}') + + + t0 = time.time() + sd_models.move_model(model, devices.device) + output = model(prompt=prompt, image=image) + json_prompt = output.values["json_prompt"] + sd_models.move_model(model, devices.cpu) + devices.torch_gc() + t1 = time.time() + log.debug(f'JSONEncode: model={model} prompt="{prompt}" json="{json_prompt}" time={t1-t0:.2f}') + return json_prompt + + +def before_prompt_encode(prompt): + if isinstance(prompt, list): + prompt = prompt[0] if len(prompt) > 0 else '' + + try: + data = json.loads(prompt) + return data + except Exception: # not a json + pass + + dct = generate_prompt_local(prompt) + return dct diff --git a/pipelines/model_bria.py b/pipelines/model_bria.py index 3f68fa774..ae9c9fa7a 100644 --- a/pipelines/model_bria.py +++ b/pipelines/model_bria.py @@ -50,6 +50,8 @@ def load_bria(checkpoint_info, diffusers_load_config=None): cache_dir=shared.opts.diffusers_dir, **load_args, ) + from pipelines.bria import prompt_to_json + pipe.before_prompt_encode = prompt_to_json.before_prompt_encode pipe.task_args = { 'output_type': 'np', }