fix ipadapter unload/reapply and use in control

This commit is contained in:
Vladimir Mandic
2024-01-30 17:02:31 -05:00
parent 8f1f538bc9
commit ccc38dde95
5 changed files with 54 additions and 64 deletions
+11 -8
View File
@@ -5,15 +5,15 @@ from typing import List, Union
import cv2
import numpy as np
from PIL import Image
from modules.control import util
from modules.control import unit
from modules.control import processors
from modules.control import util # helper functions
from modules.control import unit # control units
from modules.control import processors # image preprocessors
from modules.control.units import controlnet # lllyasviel ControlNet
from modules.control.units import xs # VisLearn ControlNet-XS
from modules.control.units import lite # Kohya ControlLLLite
from modules.control.units import t2iadapter # TencentARC T2I-Adapter
from modules.control.units import reference # ControlNet-Reference
from modules import devices, shared, errors, processing, images, sd_models, scripts, masking, ipadapter # pylint: disable=ungrouped-imports
from modules import devices, shared, errors, processing, images, sd_models, scripts, masking
debug = shared.log.trace if os.environ.get('SD_CONTROL_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -45,6 +45,9 @@ class ControlProcessing(processing.StableDiffusionProcessingImg2Img):
self.attention = 'Attention'
self.fidelity = 0.5
self.override = None
self.ip_adapter_name = None
self.ip_adapter_scale = 1.0
self.ip_adapter_image = None
def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): # abstract
pass
@@ -498,10 +501,10 @@ def control_run(units: List[unit.Unit], inputs, inits, mask, unit_type: str, is_
if hasattr(p, 'init_images') and p.init_images is None:
del p.init_images
# ip adapter
if ipadapter.apply(shared.sd_model, p, ip_adapter, ip_scale, ip_image or input_image):
original_pipeline.feature_extractor = shared.sd_model.feature_extractor
original_pipeline.image_encoder = shared.sd_model.image_encoder
# ip adapter apply is run in processing.process_images
p.ip_adapter_name = ip_adapter
p.ip_adapter_scale = ip_scale
p.ip_adapter_image = ip_image or input_image
# pipeline
output = None
+30 -47
View File
@@ -12,10 +12,9 @@ from modules import processing, shared, devices
image_encoder = None
feature_extractor = None
image_encoder_type = None
image_encoder_name = None
loaded = None
checkpoint = None
base_repo = "h94/IP-Adapter"
ADAPTERS = {
'None': 'none',
@@ -33,9 +32,9 @@ ADAPTERS = {
def unapply(pipe): # pylint: disable=arguments-differ
try:
if pipe.unet.config.encoder_hid_dim_type == 'ip_image_proj':
# shared.log.debug('IP adapter: unload attention processor')
# pipe.unet.config.encoder_hid_dim_type = None
if hasattr(pipe, 'set_ip_adapter_scale'):
pipe.set_ip_adapter_scale(0)
if hasattr(pipe, 'unet') and hasattr(pipe.unet, 'config')and pipe.unet.config.encoder_hid_dim_type == 'ip_image_proj':
pipe.unet.encoder_hid_proj = None
pipe.config.encoder_hid_dim_type = None
pipe.unet.set_default_attn_processor()
@@ -58,7 +57,7 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_name='None', sc
unapply(pipe)
return False
# init code
global loaded, checkpoint, image_encoder, image_encoder_type, image_encoder_name # pylint: disable=global-statement
global image_encoder, image_encoder_type, image_encoder_name, feature_extractor # pylint: disable=global-statement
if pipe is None:
return False
if shared.backend != shared.Backend.DIFFUSERS:
@@ -68,17 +67,13 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_name='None', sc
shared.log.error('IP adapter: no image provided')
adapter = 'none' # unload adapter if previously loaded as it will cause runtime errors
if adapter == 'none':
if hasattr(pipe, 'set_ip_adapter_scale'):
pipe.set_ip_adapter_scale(0)
if loaded is not None:
loaded = None
unapply(pipe)
return False
unapply(pipe)
if not hasattr(pipe, 'load_ip_adapter'):
import diffusers
diffusers.StableDiffusionPipeline.load_ip_adapter()
shared.log.error(f'IP adapter: pipeline not supported: {pipe.__class__.__name__}')
return False
if shared.sd_model_type != 'sd' and shared.sd_model_type != 'sdxl':
shared.log.error(f'IP adapter: unsupported model type: {shared.sd_model_type}')
return False
# which clip to use
if 'ViT' not in adapter_name:
@@ -94,45 +89,33 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_name='None', sc
shared.log.error(f'IP adapter: unknown model type: {adapter_name}')
return False
# load image encoder used by ip adapter
if getattr(pipe, 'image_encoder', None) is None or image_encoder_name != clip_repo + '/' + subfolder or image_encoder is None:
if image_encoder_type != shared.sd_model_type or checkpoint != shared.opts.sd_model_checkpoint or image_encoder_name != clip_repo + '/' + subfolder:
if shared.sd_model_type != 'sd' and shared.sd_model_type != 'sdxl':
shared.log.error(f'IP adapter: unsupported model type: {shared.sd_model_type}')
return False
try:
from transformers import CLIPVisionModelWithProjection
shared.log.debug(f'IP adapter load: image encoder="{clip_repo}/{subfolder}"')
image_encoder = CLIPVisionModelWithProjection.from_pretrained(clip_repo, 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
image_encoder_name = clip_repo + '/' + subfolder
except Exception as e:
shared.log.error(f'IP adapter: failed to load image encoder: {e}')
return
if getattr(pipe, 'feature_extractor', None) is None:
# load feature extractor used by ip adapter
if feature_extractor is None:
from transformers import CLIPImageProcessor
shared.log.debug('IP adapter load: feature extractor')
pipe.feature_extractor = CLIPImageProcessor()
feature_extractor = CLIPImageProcessor()
# load image encoder used by ip adapter
if image_encoder is None or image_encoder_name != clip_repo + '/' + subfolder or image_encoder_type != shared.sd_model_type:
try:
from transformers import CLIPVisionModelWithProjection
shared.log.debug(f'IP adapter load: image encoder="{clip_repo}/{subfolder}"')
image_encoder = CLIPVisionModelWithProjection.from_pretrained(clip_repo, 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
image_encoder_name = clip_repo + '/' + subfolder
except Exception as e:
shared.log.error(f'IP adapter: failed to load image encoder: {e}')
return
# 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 checkpoint != shared.opts.sd_model_checkpoint or pipe.image_encoder is None:
t0 = time.time()
if loaded is not None:
shared.log.debug('IP adapter: reset attention processor')
loaded = None
else:
shared.log.debug('IP adapter: load attention processor')
pipe.image_encoder = image_encoder
subfolder = 'models' if shared.sd_model_type == 'sd' else 'sdxl_models'
pipe.load_ip_adapter(base_repo, subfolder=subfolder, weight_name=adapter)
t1 = time.time()
shared.log.info(f'IP adapter load: adapter="{adapter}" scale={scale} image={image} time={t1-t0:.2f}')
loaded = adapter
checkpoint = shared.opts.sd_model_checkpoint
else:
shared.log.debug(f'IP adapter cache: adapter="{adapter}" scale={scale} image={image}')
t0 = time.time()
subfolder = 'models' if shared.sd_model_type == 'sd' else 'sdxl_models'
pipe.image_encoder = image_encoder
pipe.feature_extractor = feature_extractor
pipe.load_ip_adapter(base_repo, subfolder=subfolder, weight_name=adapter)
pipe.set_ip_adapter_scale(scale)
t1 = time.time()
shared.log.info(f'IP adapter: adapter="{adapter}" scale={scale} image={image} time={t1-t0:.2f}')
if isinstance(image, str):
from modules.api.api import decode_base64_to_image
+2 -1
View File
@@ -769,7 +769,9 @@ def process_images(p: StableDiffusionProcessing) -> Processed:
finally:
if not shared.opts.cuda_compile:
sd_models.apply_token_merging(p.sd_model, 0)
script_callbacks.after_process_callback(p)
if p.override_settings_restore_afterwards: # restore opts to original state
for k, v in stored_opts.items():
setattr(shared.opts, k, v)
@@ -1193,7 +1195,6 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
if shared.opts.sd_vae_sliced_encode and len(decoded_samples) > 1:
samples = torch.stack([self.sd_model.get_first_stage_encoding(self.sd_model.encode_first_stage(torch.unsqueeze(resized_sample, 0)))[0] for resized_sample in resized_samples])
else:
# TODO add TEASD support
samples = self.sd_model.get_first_stage_encoding(self.sd_model.encode_first_stage(resized_samples))
image_conditioning = self.img2img_image_conditioning(resized_samples, samples)
else:
+1 -1
View File
@@ -227,7 +227,7 @@ def create_ui(_blocks: gr.Blocks=None):
with gr.Row():
with gr.Column():
gr.HTML('<a href="https://github.com/tencent-ailab/IP-Adapter">IP-Adapter</a>')
ip_adapter_name = gr.Dropdown(label='Adapter', choices=ipadapter.ADAPTERS, value='none')
ip_adapter_name = 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.Column():
ip_image = gr.Image(label="Input", show_label=False, type="pil", source="upload", interactive=True, tool="editor", height=256, width=256)