Skip to content

Commit d8bebb0

Browse files
committed
Add IP Adapter support for SDXL ControlNet Inpaint pipeline
1 parent dd63168 commit d8bebb0

2 files changed

Lines changed: 40 additions & 4 deletions

File tree

src/diffusers/pipelines/controlnet/pipeline_controlnet_inpaint_sd_xl.py

Lines changed: 39 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -19,10 +19,15 @@
1919
import PIL.Image
2020
import torch
2121
import torch.nn.functional as F
22-
from transformers import CLIPTextModel, CLIPTextModelWithProjection, CLIPTokenizer
22+
from transformers import CLIPTextModel, CLIPTextModelWithProjection, CLIPTokenizer, CLIPVisionModelWithProjection
2323

2424
from ...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+
)
2631
from ...models import AutoencoderKL, ControlNetModel, UNet2DConditionModel
2732
from ...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

142147
class 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.

tests/pipelines/controlnet/test_controlnet_inpaint_sdxl.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -135,6 +135,7 @@ def get_dummy_components(self):
135135
"tokenizer": tokenizer,
136136
"text_encoder_2": text_encoder_2,
137137
"tokenizer_2": tokenizer_2,
138+
"image_encoder": None,
138139
}
139140
return components
140141

0 commit comments

Comments
 (0)