mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 01:04:32 +02:00
@@ -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: <https://github.com/Bria-AI/Fibo-Edit/blob/master/src/fibo_edit/fibo_edit_vlm.py>
|
||||
- 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)
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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
|
||||
@@ -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',
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user