mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
flux true-scale and flux ipadapters
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
+37
-13
@@ -11,12 +11,13 @@ import json
|
||||
from PIL import Image
|
||||
import diffusers
|
||||
import transformers
|
||||
from modules import processing, shared, devices, sd_models
|
||||
from modules import processing, shared, devices, sd_models, errors
|
||||
|
||||
|
||||
clip_loaded = None
|
||||
adapters_loaded = []
|
||||
CLIP_ID = "h94/IP-Adapter"
|
||||
OPEN_ID = "openai/clip-vit-large-patch14"
|
||||
SIGLIP_ID = 'google/siglip-so400m-patch14-384'
|
||||
ADAPTERS_NONE = {
|
||||
'None': { 'name': 'none', 'repo': 'none', 'subfolder': 'none' },
|
||||
@@ -136,14 +137,22 @@ def crop_images(images, crops):
|
||||
return images
|
||||
|
||||
|
||||
def unapply(pipe): # pylint: disable=arguments-differ
|
||||
def unapply(pipe, unload: bool = False): # pylint: disable=arguments-differ
|
||||
if len(adapters_loaded) == 0:
|
||||
return
|
||||
try:
|
||||
if hasattr(pipe, 'set_ip_adapter_scale'):
|
||||
pipe.set_ip_adapter_scale(0)
|
||||
pipe.unload_ip_adapter()
|
||||
if hasattr(pipe, 'unet') and hasattr(pipe.unet, 'config') and pipe.unet.config.encoder_hid_dim_type == 'ip_image_proj':
|
||||
if unload:
|
||||
shared.log.debug('IP adapter unload')
|
||||
pipe.unload_ip_adapter()
|
||||
if hasattr(pipe, 'unet'):
|
||||
module = pipe.unet
|
||||
elif hasattr(pipe, 'transformer'):
|
||||
module = pipe.transformer
|
||||
else:
|
||||
module = None
|
||||
if module is not None and hasattr(module, 'config') and module.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()
|
||||
@@ -171,8 +180,8 @@ def load_image_encoder(pipe: diffusers.DiffusionPipeline, adapter_names: list[st
|
||||
clip_repo = SIGLIP_ID
|
||||
clip_subfolder = None
|
||||
elif shared.sd_model_type == 'f1':
|
||||
shared.log.error(f'IP adapter: adapter={adapter_name} type={shared.sd_model_type} cls={shared.sd_model.__class__.__name__}: unsupported base model')
|
||||
return False
|
||||
clip_repo = OPEN_ID
|
||||
clip_subfolder = None
|
||||
else:
|
||||
shared.log.error(f'IP adapter: unknown model type: {adapter_name}')
|
||||
return False
|
||||
@@ -181,13 +190,22 @@ def load_image_encoder(pipe: diffusers.DiffusionPipeline, adapter_names: list[st
|
||||
if pipe.image_encoder is None or clip_loaded != f'{clip_repo}/{clip_subfolder}':
|
||||
try:
|
||||
if shared.sd_model_type == 'sd3':
|
||||
pipe.image_encoder = transformers.SiglipVisionModel.from_pretrained(clip_repo, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir)
|
||||
image_encoder = transformers.SiglipVisionModel.from_pretrained(clip_repo, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir)
|
||||
else:
|
||||
pipe.image_encoder = transformers.CLIPVisionModelWithProjection.from_pretrained(clip_repo, subfolder=clip_subfolder, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir, use_safetensors=True)
|
||||
shared.log.debug(f'IP adapter load: encoder="{clip_repo}/{clip_subfolder}" cls={pipe.image_encoder.__class__.__name__}')
|
||||
if clip_subfolder is None:
|
||||
image_encoder = transformers.CLIPVisionModelWithProjection.from_pretrained(clip_repo, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir, use_safetensors=True)
|
||||
shared.log.debug(f'IP adapter load: encoder="{clip_repo}" cls={pipe.image_encoder.__class__.__name__}')
|
||||
else:
|
||||
image_encoder = transformers.CLIPVisionModelWithProjection.from_pretrained(clip_repo, subfolder=clip_subfolder, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir, use_safetensors=True)
|
||||
shared.log.debug(f'IP adapter load: encoder="{clip_repo}/{clip_subfolder}" cls={pipe.image_encoder.__class__.__name__}')
|
||||
if hasattr(pipe, 'register_modules'):
|
||||
pipe.register_modules(image_encoder=image_encoder)
|
||||
else:
|
||||
pipe.image_encoder = image_encoder
|
||||
clip_loaded = f'{clip_repo}/{clip_subfolder}'
|
||||
except Exception as e:
|
||||
shared.log.error(f'IP adapter load: encoder="{clip_repo}/{clip_subfolder}" {e}')
|
||||
errors.display(e, 'IP adapter: type=encoder')
|
||||
return False
|
||||
sd_models.move_model(pipe.image_encoder, devices.device)
|
||||
return True
|
||||
@@ -198,12 +216,17 @@ def load_feature_extractor(pipe):
|
||||
if pipe.feature_extractor is None:
|
||||
try:
|
||||
if shared.sd_model_type == 'sd3':
|
||||
pipe.feature_extractor = transformers.SiglipImageProcessor.from_pretrained(SIGLIP_ID, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir)
|
||||
feature_extractor = transformers.SiglipImageProcessor.from_pretrained(SIGLIP_ID, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir)
|
||||
else:
|
||||
pipe.feature_extractor = transformers.CLIPImageProcessor()
|
||||
feature_extractor = transformers.CLIPImageProcessor()
|
||||
if hasattr(pipe, 'register_modules'):
|
||||
pipe.register_modules(feature_extractor=feature_extractor)
|
||||
else:
|
||||
pipe.feature_extractor = feature_extractor
|
||||
shared.log.debug(f'IP adapter load: extractor={pipe.feature_extractor.__class__.__name__}')
|
||||
except Exception as e:
|
||||
shared.log.error(f'IP adapter load: extractor {e}')
|
||||
errors.display(e, 'IP adapter: type=extractor')
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -266,7 +289,7 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt
|
||||
adapters = [ADAPTERS.get(adapter_name, None) for adapter_name in adapter_names if adapter_name.lower() != 'none']
|
||||
|
||||
if len(adapters) == 0:
|
||||
unapply(pipe)
|
||||
unapply(pipe, getattr(p, 'ip_adapter_unload', False))
|
||||
if hasattr(p, 'ip_adapter_images'):
|
||||
del p.ip_adapter_images
|
||||
return False
|
||||
@@ -286,7 +309,7 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt
|
||||
shared.log.error('IP adapter: no image provided')
|
||||
adapters = [] # unload adapter if previously loaded as it will cause runtime errors
|
||||
if len(adapters) == 0:
|
||||
unapply(pipe)
|
||||
unapply(pipe, getattr(p, 'ip_adapter_unload', False))
|
||||
if hasattr(p, 'ip_adapter_images'):
|
||||
del p.ip_adapter_images
|
||||
return False
|
||||
@@ -335,4 +358,5 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt
|
||||
shared.log.info(f'IP adapter: {ip_str} image={adapter_images} mask={adapter_masks is not None} time={t1-t0:.2f}')
|
||||
except Exception as e:
|
||||
shared.log.error(f'IP adapter load: adapters={adapter_names} repo={repos} folders={subfolders} names={names} {e}')
|
||||
errors.display(e, 'IP adapter: type=adapter')
|
||||
return True
|
||||
|
||||
Reference in New Issue
Block a user