1919import PIL .Image
2020import torch
2121import torch .nn .functional as F
22- from transformers import CLIPTextModel , CLIPTextModelWithProjection , CLIPTokenizer
22+ from transformers import CLIPTextModel , CLIPTextModelWithProjection , CLIPTokenizer , CLIPVisionModelWithProjection
2323
2424from ...image_processor import PipelineImageInput , VaeImageProcessor
25- from ...loaders import FromSingleFileMixin , StableDiffusionXLLoraLoaderMixin , TextualInversionLoaderMixin
25+ from ...loaders import (
26+ FromSingleFileMixin ,
27+ IPAdapterMixin ,
28+ StableDiffusionXLLoraLoaderMixin ,
29+ TextualInversionLoaderMixin ,
30+ )
2631from ...models import AutoencoderKL , ControlNetModel , UNet2DConditionModel
2732from ...models .attention_processor import (
2833 AttnProcessor2_0 ,
@@ -140,7 +145,7 @@ def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
140145
141146
142147class StableDiffusionXLControlNetInpaintPipeline (
143- DiffusionPipeline , StableDiffusionXLLoraLoaderMixin , FromSingleFileMixin
148+ DiffusionPipeline , StableDiffusionXLLoraLoaderMixin , IPAdapterMixin , FromSingleFileMixin
144149):
145150 r"""
146151 Pipeline for text-to-image generation using Stable Diffusion XL.
@@ -152,6 +157,7 @@ class StableDiffusionXLControlNetInpaintPipeline(
152157 - [`~loaders.StableDiffusionXLLoraLoaderMixin.load_lora_weights`] for loading LoRA weights
153158 - [`~loaders.StableDiffusionXLLoraLoaderMixin.save_lora_weights`] for saving LoRA weights
154159 - [`~loaders.FromSingleFileMixin.from_single_file`] for loading `.ckpt` files
160+ - [`~loaders.IPAdapterMixin.load_ip_adapter`] for loading IP Adapters
155161
156162 Args:
157163 vae ([`AutoencoderKL`]):
@@ -179,7 +185,7 @@ class StableDiffusionXLControlNetInpaintPipeline(
179185 """
180186
181187 model_cpu_offload_seq = "text_encoder->text_encoder_2->unet->vae"
182- _optional_components = ["tokenizer" , "tokenizer_2" , "text_encoder" , "text_encoder_2" ]
188+ _optional_components = ["tokenizer" , "tokenizer_2" , "text_encoder" , "text_encoder_2" , "image_encoder" ]
183189 _callback_tensor_inputs = ["latents" , "prompt_embeds" , "negative_prompt_embeds" ]
184190
185191 def __init__ (
@@ -192,6 +198,7 @@ def __init__(
192198 unet : UNet2DConditionModel ,
193199 controlnet : ControlNetModel ,
194200 scheduler : KarrasDiffusionSchedulers ,
201+ image_encoder : CLIPVisionModelWithProjection = None ,
195202 requires_aesthetics_score : bool = False ,
196203 force_zeros_for_empty_prompt : bool = True ,
197204 add_watermarker : Optional [bool ] = None ,
@@ -210,6 +217,7 @@ def __init__(
210217 unet = unet ,
211218 controlnet = controlnet ,
212219 scheduler = scheduler ,
220+ image_encoder = image_encoder ,
213221 )
214222 self .register_to_config (force_zeros_for_empty_prompt = force_zeros_for_empty_prompt )
215223 self .register_to_config (requires_aesthetics_score = requires_aesthetics_score )
@@ -497,6 +505,22 @@ def encode_prompt(
497505
498506 return prompt_embeds , negative_prompt_embeds , pooled_prompt_embeds , negative_pooled_prompt_embeds
499507
508+ # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.encode_image
509+ def encode_image (self , image , device , num_images_per_prompt ):
510+ dtype = next (self .image_encoder .parameters ()).dtype
511+
512+ if not isinstance (image , torch .Tensor ):
513+ image = self .feature_extractor (
514+ image , return_tensors = "pt" ).pixel_values
515+
516+ image = image .to (device = device , dtype = dtype )
517+ image_embeds = self .image_encoder (image ).image_embeds
518+ image_embeds = image_embeds .repeat_interleave (
519+ num_images_per_prompt , dim = 0 )
520+
521+ uncond_image_embeds = torch .zeros_like (image_embeds )
522+ return image_embeds , uncond_image_embeds
523+
500524 # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_extra_step_kwargs
501525 def prepare_extra_step_kwargs (self , generator , eta ):
502526 # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
@@ -1079,6 +1103,7 @@ def __call__(
10791103 latents : Optional [torch .FloatTensor ] = None ,
10801104 prompt_embeds : Optional [torch .FloatTensor ] = None ,
10811105 negative_prompt_embeds : Optional [torch .FloatTensor ] = None ,
1106+ ip_adapter_image : Optional [PipelineImageInput ] = None ,
10821107 pooled_prompt_embeds : Optional [torch .FloatTensor ] = None ,
10831108 negative_pooled_prompt_embeds : Optional [torch .FloatTensor ] = None ,
10841109 output_type : Optional [str ] = "pil" ,
@@ -1167,6 +1192,7 @@ def __call__(
11671192 Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
11681193 weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
11691194 argument.
1195+ ip_adapter_image: (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters.
11701196 pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
11711197 Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
11721198 If not provided, pooled text embeddings will be generated from `prompt` input argument.
@@ -1348,6 +1374,11 @@ def __call__(
13481374 clip_skip = self .clip_skip ,
13491375 )
13501376
1377+ if ip_adapter_image is not None :
1378+ image_embeds , negative_image_embeds = self .encode_image (ip_adapter_image , device , num_images_per_prompt )
1379+ if self .do_classifier_free_guidance :
1380+ image_embeds = torch .cat ([negative_image_embeds , image_embeds ])
1381+
13511382 # 4. set timesteps
13521383 def denoising_value_valid (dnv ):
13531384 return isinstance (denoising_end , float ) and 0 < dnv < 1
@@ -1557,6 +1588,10 @@ def denoising_value_valid(dnv):
15571588
15581589 added_cond_kwargs = {"text_embeds" : add_text_embeds , "time_ids" : add_time_ids }
15591590
1591+ # Add image embeds for IP-Adapter
1592+ if ip_adapter_image :
1593+ added_cond_kwargs ["image_embeds" ] = image_embeds
1594+
15601595 # controlnet(s) inference
15611596 if guess_mode and self .do_classifier_free_guidance :
15621597 # Infer ControlNet only for the conditional batch.
0 commit comments