photomaker with offloading

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-01-30 08:38:03 -05:00
parent 9a5a5536f6
commit 04ad0caa0c
4 changed files with 14 additions and 16 deletions
+2 -1
View File
@@ -89,7 +89,8 @@ class Script(scripts.Script):
gr.HTML('<a href="https://photo-maker.github.io/" target="_blank">&nbsp 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():
+8 -13
View File
@@ -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)
+2 -1
View File
@@ -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