mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 01:04:32 +02:00
fix ipadapter unload/reapply and use in control
This commit is contained in:
+11
-8
@@ -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
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user