mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 01:04:32 +02:00
photomaker with offloading
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
@@ -89,7 +89,8 @@ class Script(scripts.Script):
|
||||
gr.HTML('<a href="https://photo-maker.github.io/" target="_blank">  Tenecent ARC Lab PhotoMaker</a><br>')
|
||||
with gr.Row():
|
||||
pm_model = gr.Dropdown(label='PhotoMaker Model', choices=['PhotoMaker v1', 'PhotoMaker v2'], value='PhotoMaker v2')
|
||||
pm_trigger = gr.Text(label='Trigger word', value="person")
|
||||
pm_trigger = gr.Text(label='Trigger word', placeholder="enter one word in prompt")
|
||||
with gr.Row():
|
||||
pm_strength = gr.Slider(label='Strength', minimum=0.0, maximum=2.0, step=0.01, value=1.0)
|
||||
pm_start = gr.Slider(label='Start', minimum=0.0, maximum=1.0, step=0.01, value=0.5)
|
||||
with gr.Row():
|
||||
|
||||
@@ -13,6 +13,10 @@ def photo_maker(p: processing.StableDiffusionProcessing, app, model: str, input_
|
||||
shared.log.warning('PhotoMaker: no input images')
|
||||
return None
|
||||
|
||||
if len(trigger) == 0:
|
||||
shared.log.warning('PhotoMaker: no trigger word')
|
||||
return None
|
||||
|
||||
c = shared.sd_model.__class__.__name__ if shared.sd_loaded else ''
|
||||
if c != 'StableDiffusionXLPipeline':
|
||||
shared.log.warning(f'PhotoMaker invalid base model: current={c} required=StableDiffusionXLPipeline')
|
||||
@@ -35,20 +39,10 @@ def photo_maker(p: processing.StableDiffusionProcessing, app, model: str, input_
|
||||
|
||||
# create new pipeline
|
||||
orig_pipeline = shared.sd_model # backup current pipeline definition
|
||||
shared.sd_model = PhotoMakerStableDiffusionXLPipeline(
|
||||
vae = shared.sd_model.vae,
|
||||
text_encoder=shared.sd_model.text_encoder,
|
||||
text_encoder_2=shared.sd_model.text_encoder_2,
|
||||
tokenizer=shared.sd_model.tokenizer,
|
||||
tokenizer_2=shared.sd_model.tokenizer_2,
|
||||
unet=shared.sd_model.unet,
|
||||
scheduler=shared.sd_model.scheduler,
|
||||
force_zeros_for_empty_prompt=shared.opts.diffusers_force_zeros,
|
||||
)
|
||||
shared.sd_model = sd_models.switch_pipe(PhotoMakerStableDiffusionXLPipeline, shared.sd_model)
|
||||
sd_models.copy_diffuser_options(shared.sd_model, orig_pipeline) # copy options from original pipeline
|
||||
sd_models.set_diffuser_options(shared.sd_model) # set all model options such as fp16, offload, etc.
|
||||
sd_models.move_model(shared.sd_model, devices.device) # move pipeline to device
|
||||
shared.sd_model.to(dtype=devices.dtype)
|
||||
sd_models.apply_balanced_offload(shared.sd_model) # apply balanced offload
|
||||
|
||||
orig_prompt_attention = shared.opts.prompt_attention
|
||||
shared.opts.data['prompt_attention'] = 'fixed' # otherwise need to deal with class_tokens_mask
|
||||
@@ -71,6 +65,7 @@ def photo_maker(p: processing.StableDiffusionProcessing, app, model: str, input_
|
||||
trigger_word=trigger,
|
||||
weight_name='photomaker-v2.bin' if is_v2 else 'photomaker-v1.bin',
|
||||
pm_version='v2' if is_v2 else 'v1',
|
||||
device=devices.device,
|
||||
cache_dir=shared.opts.hfcache_dir,
|
||||
)
|
||||
shared.sd_model.set_adapters(["photomaker"], adapter_weights=[strength])
|
||||
@@ -83,7 +78,7 @@ def photo_maker(p: processing.StableDiffusionProcessing, app, model: str, input_
|
||||
face = sorted(faces, key=lambda x:(x['bbox'][2]-x['bbox'][0])*x['bbox'][3]-x['bbox'][1])[-1] # only use the maximum face
|
||||
id_embed_list.append(torch.from_numpy(face['embedding']))
|
||||
shared.log.debug(f'PhotoMaker: face={i+1} score={face.det_score:.2f} gender={"female" if face.gender==0 else "male"} age={face.age} bbox={face.bbox}')
|
||||
p.task_args['id_embeds'] = torch.stack(id_embed_list)
|
||||
p.task_args['id_embeds'] = torch.stack(id_embed_list).to(device=devices.device, dtype=devices.dtype)
|
||||
|
||||
# run processing
|
||||
processed: processing.Processed = processing.process_images(p)
|
||||
|
||||
@@ -115,6 +115,7 @@ class PhotoMakerStableDiffusionXLPipeline(StableDiffusionXLPipeline):
|
||||
subfolder: str = '',
|
||||
trigger_word: str = 'img',
|
||||
pm_version: str = 'v2',
|
||||
device: torch.device = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -197,7 +198,7 @@ class PhotoMakerStableDiffusionXLPipeline(StableDiffusionXLPipeline):
|
||||
raise NotImplementedError(f"The PhotoMaker version [{pm_version}] does not support")
|
||||
|
||||
id_encoder.load_state_dict(state_dict["id_encoder"], strict=True)
|
||||
id_encoder = id_encoder.to(self.device, dtype=self.unet.dtype)
|
||||
id_encoder = id_encoder.to(device, dtype=self.unet.dtype)
|
||||
self.id_encoder = id_encoder # pylint: disable=attribute-defined-outside-init
|
||||
|
||||
# load lora into models
|
||||
|
||||
Reference in New Issue
Block a user