From 1d383588999562caa71ce35269fabd98362e0267 Mon Sep 17 00:00:00 2001
From: Vladimir Mandic
Date: Sat, 23 Dec 2023 12:50:30 -0500
Subject: [PATCH] control add ip-adapter
---
modules/control/adapters.py | 5 +-
modules/control/controlnets.py | 5 +-
modules/control/controlnetsxs.py | 5 +-
modules/control/ipadapter.py | 90 ++++++++++++++++++++++++++++++++
modules/control/reference.py | 5 +-
modules/control/run.py | 18 +++++--
modules/processing.py | 4 +-
modules/sd_models.py | 4 ++
modules/ui_control.py | 28 ++++++++--
wiki | 2 +-
10 files changed, 150 insertions(+), 16 deletions(-)
create mode 100644 modules/control/ipadapter.py
diff --git a/modules/control/adapters.py b/modules/control/adapters.py
index 12ed28907..19030d2cb 100644
--- a/modules/control/adapters.py
+++ b/modules/control/adapters.py
@@ -124,6 +124,8 @@ class AdapterPipeline():
tokenizer_2=pipeline.tokenizer_2,
unet=pipeline.unet,
scheduler=pipeline.scheduler,
+ image_encoder=getattr(pipeline, 'image_encoder', None),
+ feature_extractor=getattr(pipeline, 'feature_extractor', None),
adapter=adapter,
).to(pipeline.device)
elif isinstance(pipeline, StableDiffusionPipeline):
@@ -133,9 +135,10 @@ class AdapterPipeline():
tokenizer=pipeline.tokenizer,
unet=pipeline.unet,
scheduler=pipeline.scheduler,
+ image_encoder=getattr(pipeline, 'image_encoder', None),
+ feature_extractor=getattr(pipeline, 'feature_extractor', None),
requires_safety_checker=False,
safety_checker=None,
- feature_extractor=None,
adapter=adapter,
).to(pipeline.device)
else:
diff --git a/modules/control/controlnets.py b/modules/control/controlnets.py
index 0230fe773..0e9816082 100644
--- a/modules/control/controlnets.py
+++ b/modules/control/controlnets.py
@@ -139,6 +139,8 @@ class ControlNetPipeline():
tokenizer_2=pipeline.tokenizer_2,
unet=pipeline.unet,
scheduler=pipeline.scheduler,
+ image_encoder=getattr(pipeline, 'image_encoder', None),
+ feature_extractor=getattr(pipeline, 'feature_extractor', None),
controlnet=controlnet, # can be a list
).to(pipeline.device)
elif isinstance(pipeline, StableDiffusionPipeline):
@@ -148,9 +150,10 @@ class ControlNetPipeline():
tokenizer=pipeline.tokenizer,
unet=pipeline.unet,
scheduler=pipeline.scheduler,
+ image_encoder=getattr(pipeline, 'image_encoder', None),
+ feature_extractor=getattr(pipeline, 'feature_extractor', None),
requires_safety_checker=False,
safety_checker=None,
- feature_extractor=None,
controlnet=controlnet, # can be a list
).to(pipeline.device)
else:
diff --git a/modules/control/controlnetsxs.py b/modules/control/controlnetsxs.py
index 7096902a7..630f06df3 100644
--- a/modules/control/controlnetsxs.py
+++ b/modules/control/controlnetsxs.py
@@ -131,6 +131,8 @@ class ControlNetXSPipeline():
tokenizer_2=pipeline.tokenizer_2,
unet=pipeline.unet,
scheduler=pipeline.scheduler,
+ image_encoder=getattr(pipeline, 'image_encoder', None),
+ feature_extractor=getattr(pipeline, 'feature_extractor', None),
controlnet=controlnet, # can be a list
).to(pipeline.device)
elif isinstance(pipeline, StableDiffusionPipeline):
@@ -140,9 +142,10 @@ class ControlNetXSPipeline():
tokenizer=pipeline.tokenizer,
unet=pipeline.unet,
scheduler=pipeline.scheduler,
+ image_encoder=getattr(pipeline, 'image_encoder', None),
+ feature_extractor=getattr(pipeline, 'feature_extractor', None),
requires_safety_checker=False,
safety_checker=None,
- feature_extractor=None,
controlnet=controlnet, # can be a list
).to(pipeline.device)
else:
diff --git a/modules/control/ipadapter.py b/modules/control/ipadapter.py
new file mode 100644
index 000000000..ed12ca63c
--- /dev/null
+++ b/modules/control/ipadapter.py
@@ -0,0 +1,90 @@
+import time
+from PIL import Image
+from modules import shared, processing, devices
+
+
+image_encoder = None
+image_encoder_type = None
+loaded = None
+ADAPTERS = [
+ 'none',
+ 'ip-adapter_sd15',
+ 'ip-adapter_sd15_light',
+ 'ip-adapter-plus_sd15',
+ 'ip-adapter-plus-face_sd15',
+ 'ip-adapter-full-face_sd15',
+ # 'models/ip-adapter_sd15_vit-G', # RuntimeError: mat1 and mat2 shapes cannot be multiplied (2x1024 and 1280x3072)
+ 'ip-adapter_sdxl',
+ # 'sdxl_models/ip-adapter_sdxl_vit-h',
+ # 'sdxl_models/ip-adapter-plus_sdxl_vit-h',
+ # 'sdxl_models/ip-adapter-plus-face_sdxl_vit-h',
+]
+
+
+def apply_ip_adapter(pipe, p: processing.StableDiffusionProcessing, adapter, scale, image, reset=False): # pylint: disable=arguments-differ
+ from transformers import CLIPVisionModelWithProjection
+ # overrides
+ if hasattr(p, 'ip_adapter_name'):
+ adapter = p.ip_adapter_name
+ if hasattr(p, 'ip_adapter_scale'):
+ scale = p.ip_adapter_scale
+ if hasattr(p, 'ip_adapter_image'):
+ image = p.ip_adapter_image
+ # init code
+ global loaded, image_encoder, image_encoder_type # pylint: disable=global-statement
+ if pipe is None:
+ return
+ if shared.backend != shared.Backend.DIFFUSERS:
+ shared.log.warning('IP adapter: not in diffusers mode')
+ return False
+ if adapter == 'none':
+ if hasattr(pipe, 'set_ip_adapter_scale'):
+ pipe.set_ip_adapter_scale(0)
+ if loaded is not None:
+ shared.log.debug('IP adapter: unload attention processor')
+ pipe.unet.set_default_attn_processor()
+ pipe.unet.config.encoder_hid_dim_type = None
+ loaded = None
+ return False
+ if image is None:
+ image = Image.new('RGB', (512, 512), (0, 0, 0))
+ if not hasattr(pipe, 'load_ip_adapter'):
+ shared.log.error(f'IP adapter: pipeline not supported: {pipe.__class__.__name__}')
+ return False
+ if getattr(pipe, 'image_encoder', None) is None or getattr(pipe, 'image_encoder', None) == (None, None):
+ if shared.sd_model_type == 'sd':
+ subfolder = 'models/image_encoder'
+ elif shared.sd_model_type == 'sdxl':
+ subfolder = 'sdxl_models/image_encoder'
+ else:
+ shared.log.error(f'IP adapter: unsupported model type: {shared.sd_model_type}')
+ return False
+ if image_encoder is None or image_encoder_type != shared.sd_model_type:
+ try:
+ image_encoder = CLIPVisionModelWithProjection.from_pretrained("h94/IP-Adapter", subfolder=subfolder, torch_dtype=devices.dtype, cache_dir=shared.opts.diffusers_dir, use_safetensors=True).to(devices.device)
+ image_encoder_type = shared.sd_model_type
+ except Exception as e:
+ shared.log.error(f'IP adapter: failed to load image encoder: {e}')
+ return False
+ pipe.image_encoder = image_encoder
+
+ # main code
+ subfolder = 'models' if 'sd15' in adapter else 'sdxl_models'
+ if adapter != loaded or getattr(pipe.unet.config, 'encoder_hid_dim_type', None) is None or reset:
+ t0 = time.time()
+ if loaded is not None:
+ # shared.log.debug('IP adapter: reset attention processor')
+ pipe.unet.set_default_attn_processor()
+ loaded = None
+ else:
+ shared.log.debug('IP adapter: load attention processor')
+ pipe.load_ip_adapter("h94/IP-Adapter", subfolder=subfolder, weight_name=f'{adapter}.safetensors')
+ t1 = time.time()
+ shared.log.info(f'IP adapter load: adapter="{adapter}" scale={scale} image={image} time={t1-t0:.2f}')
+ loaded = adapter
+ else:
+ shared.log.debug(f'IP adapter cache: adapter="{adapter}" scale={scale} image={image}')
+ pipe.set_ip_adapter_scale(scale)
+ p.task_args['ip_adapter_image'] = p.batch_size * [image]
+ p.extra_generation_params["IP Adapter"] = f'{adapter}:{scale}'
+ return True
diff --git a/modules/control/reference.py b/modules/control/reference.py
index 8116274a7..32d8dcdd2 100644
--- a/modules/control/reference.py
+++ b/modules/control/reference.py
@@ -29,6 +29,8 @@ class ReferencePipeline():
tokenizer_2=pipeline.tokenizer_2,
unet=pipeline.unet,
scheduler=pipeline.scheduler,
+ image_encoder=getattr(pipeline, 'image_encoder', None),
+ feature_extractor=getattr(pipeline, 'feature_extractor', None),
).to(pipeline.device)
elif isinstance(pipeline, StableDiffusionPipeline):
self.pipeline = StableDiffusionReferencePipeline(
@@ -37,9 +39,10 @@ class ReferencePipeline():
tokenizer=pipeline.tokenizer,
unet=pipeline.unet,
scheduler=pipeline.scheduler,
+ image_encoder=getattr(pipeline, 'image_encoder', None),
+ feature_extractor=getattr(pipeline, 'feature_extractor', None),
requires_safety_checker=False,
safety_checker=None,
- feature_extractor=None,
).to(pipeline.device)
else:
log.error(f'Control {what} pipeline: class={pipeline.__class__.__name__} unsupported model type')
diff --git a/modules/control/run.py b/modules/control/run.py
index e2dad2ea6..a80829e28 100644
--- a/modules/control/run.py
+++ b/modules/control/run.py
@@ -13,6 +13,7 @@ from modules.control import controlnets # lllyasviel ControlNet
from modules.control import controlnetsxs # VisLearn ControlNet-XS
from modules.control import adapters # TencentARC T2I-Adapter
from modules.control import reference # ControlNet-Reference
+from modules.control import ipadapter # IP-Adapter
from modules import devices, shared, errors, processing, images, sd_models, sd_samplers
@@ -67,6 +68,7 @@ def control_run(units: List[unit.Unit], inputs, inits, unit_type: str, is_genera
resize_mode, resize_name, width, height, scale_by, selected_scale_tab, resize_time,
denoising_strength, batch_count, batch_size,
video_skip_frames, video_type, video_duration, video_loop, video_pad, video_interpolate,
+ ip_adapter, ip_scale, ip_image, ip_type,
):
global pipe, original_pipeline # pylint: disable=global-statement
debug(f'Control {unit_type}: input={inputs} init={inits} type={input_type}')
@@ -246,6 +248,9 @@ def control_run(units: List[unit.Unit], inputs, inits, unit_type: str, is_genera
shared.sd_model.to(shared.device)
sd_models.copy_diffuser_options(shared.sd_model, original_pipeline) # copy options from original pipeline
sd_models.set_diffuser_options(shared.sd_model)
+ if ipadapter.apply_ip_adapter(shared.sd_model, p, ip_adapter, ip_scale, ip_image, reset=True):
+ original_pipeline.feature_extractor = shared.sd_model.feature_extractor
+ original_pipeline.image_encoder = shared.sd_model.image_encoder
try:
with devices.inference_context():
@@ -377,6 +382,9 @@ def control_run(units: List[unit.Unit], inputs, inits, unit_type: str, is_genera
else:
p.task_args['image'] = init_image
+ if ip_type == 1 and ip_adapter != 'none':
+ p.task_args['ip_adapter_image'] = input_image
+
if is_generator:
image_txt = f'{processed_image.width}x{processed_image.height}' if processed_image is not None else 'None'
msg = f'process | {index} of {frames if video is not None else len(inputs)} | {"Image" if video is None else "Frame"} {image_txt}'
@@ -391,18 +399,18 @@ def control_run(units: List[unit.Unit], inputs, inits, unit_type: str, is_genera
if not has_models and (unit_type == 'controlnet' or unit_type == 'adapter' or unit_type == 'xs'): # run in txt2img or img2img mode
if processed_image is not None:
p.init_images = [processed_image]
- pipe = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.IMAGE_2_IMAGE)
+ shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.IMAGE_2_IMAGE)
else:
- pipe = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.TEXT_2_IMAGE)
+ shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.TEXT_2_IMAGE)
elif unit_type == 'reference':
p.is_control = True
- pipe = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.TEXT_2_IMAGE)
+ shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.TEXT_2_IMAGE)
else: # actual control
p.is_control = True
if 'control_image' in p.task_args:
- pipe = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.IMAGE_2_IMAGE) # only controlnet supports img2img
+ shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.IMAGE_2_IMAGE) # only controlnet supports img2img
else:
- pipe = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.TEXT_2_IMAGE)
+ shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.TEXT_2_IMAGE)
# pipeline
output = None
diff --git a/modules/processing.py b/modules/processing.py
index b2ea8c175..e2283b8ae 100644
--- a/modules/processing.py
+++ b/modules/processing.py
@@ -1261,9 +1261,9 @@ class StableDiffusionProcessingImg2Img(StableDiffusionProcessing):
self.mask_blur_y = value
def init(self, all_prompts, all_seeds, all_subseeds):
- if shared.backend == shared.Backend.DIFFUSERS and self.image_mask is not None:
+ if shared.backend == shared.Backend.DIFFUSERS and self.image_mask is not None and not self.is_control:
shared.sd_model = modules.sd_models.set_diffuser_pipe(self.sd_model, modules.sd_models.DiffusersTaskType.INPAINTING)
- elif shared.backend == shared.Backend.DIFFUSERS and self.image_mask is None:
+ elif shared.backend == shared.Backend.DIFFUSERS and self.image_mask is None and not self.is_control:
shared.sd_model = modules.sd_models.set_diffuser_pipe(self.sd_model, modules.sd_models.DiffusersTaskType.IMAGE_2_IMAGE)
if self.sampler_name == "PLMS":
diff --git a/modules/sd_models.py b/modules/sd_models.py
index 6fe672b4f..5d0dce124 100644
--- a/modules/sd_models.py
+++ b/modules/sd_models.py
@@ -1027,6 +1027,8 @@ def set_diffuser_pipe(pipe, new_pipe_type):
sd_model_hash = getattr(pipe, "sd_model_hash", None)
has_accelerate = getattr(pipe, "has_accelerate", None)
embedding_db = getattr(pipe, "embedding_db", None)
+ image_encoder = getattr(pipe, "image_encoder", None)
+ feature_extractor = getattr(pipe, "feature_extractor", None)
# TODO implement alternative diffusion pipelines
"""
@@ -1065,6 +1067,8 @@ def set_diffuser_pipe(pipe, new_pipe_type):
new_pipe.sd_model_hash = sd_model_hash
new_pipe.has_accelerate = has_accelerate
new_pipe.embedding_db = embedding_db
+ new_pipe.image_encoder = image_encoder
+ new_pipe.feature_extractor = feature_extractor
new_pipe.is_sdxl = True # pylint: disable=attribute-defined-outside-init # a1111 compatibility item
new_pipe.is_sd2 = False # pylint: disable=attribute-defined-outside-init
new_pipe.is_sd1 = False # pylint: disable=attribute-defined-outside-init
diff --git a/modules/ui_control.py b/modules/ui_control.py
index 7665843e0..607a78f7e 100644
--- a/modules/ui_control.py
+++ b/modules/ui_control.py
@@ -6,6 +6,7 @@ from modules.control import controlnetsxs # vislearn ControlNet-XS
from modules.control import adapters # TencentARC T2I-Adapter
from modules.control import processors # patrickvonplaten controlnet_aux
from modules.control import reference # reference pipeline
+from modules.control import ipadapter # reference pipeline
from modules import errors, shared, progress, sd_samplers, ui, ui_components, ui_symbols, ui_common, generation_parameters_copypaste, call_queue
from modules.ui_components import FormRow, FormGroup
@@ -202,9 +203,14 @@ def create_ui(_blocks: gr.Blocks=None):
with gr.Row(elem_id='control_settings'):
with gr.Accordion(open=False, label="Input", elem_id="control_input", elem_classes=["small-accordion"]):
- input_type = gr.Radio(label="Input type", choices=['Control only', 'Init image same as control', 'Separate init image'], value='Control only', type='index', elem_id='control_input_type')
- denoising_strength = gr.Slider(minimum=0.01, maximum=0.99, step=0.01, label='Denoising strength', value=0.50, elem_id="control_denoising_strength")
- show_preview = gr.Checkbox(label="Show preview", value=True, elem_id="control_show_preview")
+ with gr.Row():
+ input_type = gr.Radio(label="Input type", choices=['Control only', 'Init image same as control', 'Separate init image'], value='Control only', type='index', elem_id='control_input_type')
+ with gr.Row():
+ denoising_strength = gr.Slider(minimum=0.01, maximum=0.99, step=0.01, label='Denoising strength', value=0.50, elem_id="control_denoising_strength")
+ with gr.Row():
+ show_ip = gr.Checkbox(label="Enable IP adapter", value=False, elem_id="control_show_ip")
+ with gr.Row():
+ show_preview = gr.Checkbox(label="Show preview", value=False, elem_id="control_show_preview")
resize_mode, resize_name, width, height, scale_by, selected_scale_tab, resize_time = ui.create_resize_inputs('control', [], time_selector=True, scale_visible=False, mode='Fixed')
@@ -261,6 +267,17 @@ def create_ui(_blocks: gr.Blocks=None):
init_batch = gr.File(label="Input", show_label=False, file_count='multiple', file_types=['image'], type='file', interactive=True, height=gr_height)
with gr.Tab('Folder', id='init-folder') as tab_folder_init:
init_folder = gr.File(label="Input", show_label=False, file_count='directory', file_types=['image'], type='file', interactive=True, height=gr_height)
+ with gr.Column(scale=9, elem_id='control-init-column', visible=False) as column_ip:
+ gr.HTML('IP Adapter
')
+ with gr.Tabs(elem_classes=['control-tabs'], elem_id='control-tab-ip'):
+ with gr.Tab('Image', id='init-image') as tab_image_init:
+ ip_image = gr.Image(label="Input", show_label=False, type="pil", source="upload", interactive=True, tool="editor", height=gr_height)
+ with gr.Row():
+ ip_adapter = gr.Dropdown(label='Adapter', choices=ipadapter.ADAPTERS, value='none')
+ ip_scale = gr.Slider(label='Scale', minimum=0.0, maximum=1.0, step=0.01, value=0.5)
+ with gr.Row():
+ ip_type = gr.Radio(label="Input type", choices=['Init image same as control', 'Separate init image'], value='Init image same as control', type='index', elem_id='control_ip_type')
+ ip_image.change(fn=lambda x: gr.update(value='Init image same as control' if x is None else 'Separate init image'), inputs=[ip_image], outputs=[ip_type])
with gr.Column(scale=9, elem_id='control-output-column', visible=True) as _column_output:
gr.HTML('Output')
with gr.Tabs(elem_classes=['control-tabs'], elem_id='control-tab-output') as output_tabs:
@@ -270,7 +287,7 @@ def create_ui(_blocks: gr.Blocks=None):
output_image = gr.Image(label="Input", show_label=False, type="pil", interactive=False, tool="editor", height=gr_height)
with gr.Tab('Video', id='out-video'):
output_video = gr.Video(label="Input", show_label=False, height=gr_height)
- with gr.Column(scale=9, elem_id='control-preview-column', visible=True) as column_preview:
+ with gr.Column(scale=9, elem_id='control-preview-column', visible=False) as column_preview:
gr.HTML('Preview')
with gr.Tabs(elem_classes=['control-tabs'], elem_id='control-tab-preview'):
with gr.Tab('Preview', id='preview-image') as tab_image:
@@ -284,6 +301,7 @@ def create_ui(_blocks: gr.Blocks=None):
if hasattr(ctrl, 'select'):
ctrl.select(fn=select_input, inputs=inputs, outputs=outputs)
show_preview.change(fn=lambda x: gr.update(visible=x), inputs=[show_preview], outputs=[column_preview])
+ show_ip.change(fn=lambda x: gr.update(visible=x), inputs=[show_ip], outputs=[column_ip])
input_type.change(fn=lambda x: gr.update(visible=x == 2), inputs=[input_type], outputs=[column_init])
with gr.Tabs(elem_id='control-tabs') as _tabs_control_type:
@@ -508,6 +526,7 @@ def create_ui(_blocks: gr.Blocks=None):
resize_mode, resize_name, width, height, scale_by, selected_scale_tab, resize_time,
denoising_strength, batch_count, batch_size,
video_skip_frames, video_type, video_duration, video_loop, video_pad, video_interpolate,
+ ip_adapter, ip_scale, ip_image, ip_type,
]
output_fields = [
preview_process,
@@ -517,6 +536,7 @@ def create_ui(_blocks: gr.Blocks=None):
result_txt,
]
paste_fields = [] # TODO paste fields
+
control_dict = dict(
fn=generate_click,
_js="submit_control",
diff --git a/wiki b/wiki
index 62c496b6c..2b232b2ca 160000
--- a/wiki
+++ b/wiki
@@ -1 +1 @@
-Subproject commit 62c496b6cf936eec933d58d912c47372212c1281
+Subproject commit 2b232b2cab058b969acdea7c682252dbe4ca6cf1