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