CheckSim commited on
Commit
3ef6ae3
·
verified ·
1 Parent(s): bef8d98

Upload 2 files

Browse files
Files changed (2) hide show
  1. app.py +0 -7
  2. pipeline_stable_diffusion_xl_instantid.py +1066 -0
app.py CHANGED
@@ -11,13 +11,6 @@ from huggingface_hub import hf_hub_download, snapshot_download
11
  from insightface.app import FaceAnalysis
12
  from diffusers import ControlNetModel
13
 
14
- import urllib.request
15
- # Scarichiamo la pipeline custom di InstantID direttamente dal repository ufficiale di diffusers
16
- pipeline_url = "https://raw.githubusercontent.com/huggingface/diffusers/main/examples/community/pipeline_stable_diffusion_xl_instantid.py"
17
- if not os.path.exists("pipeline_stable_diffusion_xl_instantid.py"):
18
- print("Scaricamento pipeline InstantID custom...")
19
- urllib.request.urlretrieve(pipeline_url, "pipeline_stable_diffusion_xl_instantid.py")
20
-
21
  from pipeline_stable_diffusion_xl_instantid import StableDiffusionXLInstantIDPipeline
22
 
23
  # ==============================================================================
 
11
  from insightface.app import FaceAnalysis
12
  from diffusers import ControlNetModel
13
 
 
 
 
 
 
 
 
14
  from pipeline_stable_diffusion_xl_instantid import StableDiffusionXLInstantIDPipeline
15
 
16
  # ==============================================================================
pipeline_stable_diffusion_xl_instantid.py ADDED
@@ -0,0 +1,1066 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2025 The InstantX Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+
16
+ import math
17
+ from typing import Any, Callable, Dict, List, Optional, Tuple, Union
18
+
19
+ import cv2
20
+ import numpy as np
21
+ import PIL.Image
22
+ import torch
23
+ import torch.nn as nn
24
+
25
+ from diffusers import StableDiffusionXLControlNetPipeline
26
+ from diffusers.image_processor import PipelineImageInput
27
+ from diffusers.models import ControlNetModel
28
+ from diffusers.pipelines.controlnet.multicontrolnet import MultiControlNetModel
29
+ from diffusers.pipelines.stable_diffusion_xl import StableDiffusionXLPipelineOutput
30
+ from diffusers.utils import (
31
+ deprecate,
32
+ logging,
33
+ replace_example_docstring,
34
+ )
35
+ from diffusers.utils.import_utils import is_xformers_available
36
+ from diffusers.utils.torch_utils import is_compiled_module, is_torch_version
37
+
38
+
39
+ try:
40
+ import xformers
41
+ import xformers.ops
42
+
43
+ xformers_available = True
44
+ except Exception:
45
+ xformers_available = False
46
+
47
+ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
48
+
49
+ logger.warning(
50
+ "To use instant id pipelines, please make sure you have the `insightface` library installed: `pip install insightface`."
51
+ "Please refer to: https://huggingface.co/InstantX/InstantID for further instructions regarding inference"
52
+ )
53
+
54
+
55
+ def FeedForward(dim, mult=4):
56
+ inner_dim = int(dim * mult)
57
+ return nn.Sequential(
58
+ nn.LayerNorm(dim),
59
+ nn.Linear(dim, inner_dim, bias=False),
60
+ nn.GELU(),
61
+ nn.Linear(inner_dim, dim, bias=False),
62
+ )
63
+
64
+
65
+ def reshape_tensor(x, heads):
66
+ bs, length, width = x.shape
67
+ # (bs, length, width) --> (bs, length, n_heads, dim_per_head)
68
+ x = x.view(bs, length, heads, -1)
69
+ # (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
70
+ x = x.transpose(1, 2)
71
+ # (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
72
+ x = x.reshape(bs, heads, length, -1)
73
+ return x
74
+
75
+
76
+ class PerceiverAttention(nn.Module):
77
+ def __init__(self, *, dim, dim_head=64, heads=8):
78
+ super().__init__()
79
+ self.scale = dim_head**-0.5
80
+ self.dim_head = dim_head
81
+ self.heads = heads
82
+ inner_dim = dim_head * heads
83
+
84
+ self.norm1 = nn.LayerNorm(dim)
85
+ self.norm2 = nn.LayerNorm(dim)
86
+
87
+ self.to_q = nn.Linear(dim, inner_dim, bias=False)
88
+ self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
89
+ self.to_out = nn.Linear(inner_dim, dim, bias=False)
90
+
91
+ def forward(self, x, latents):
92
+ """
93
+ Args:
94
+ x (torch.Tensor): image features
95
+ shape (b, n1, D)
96
+ latent (torch.Tensor): latent features
97
+ shape (b, n2, D)
98
+ """
99
+ x = self.norm1(x)
100
+ latents = self.norm2(latents)
101
+
102
+ b, l, _ = latents.shape
103
+
104
+ q = self.to_q(latents)
105
+ kv_input = torch.cat((x, latents), dim=-2)
106
+ k, v = self.to_kv(kv_input).chunk(2, dim=-1)
107
+
108
+ q = reshape_tensor(q, self.heads)
109
+ k = reshape_tensor(k, self.heads)
110
+ v = reshape_tensor(v, self.heads)
111
+
112
+ # attention
113
+ scale = 1 / math.sqrt(math.sqrt(self.dim_head))
114
+ weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
115
+ weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
116
+ out = weight @ v
117
+
118
+ out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
119
+
120
+ return self.to_out(out)
121
+
122
+
123
+ class Resampler(nn.Module):
124
+ def __init__(
125
+ self,
126
+ dim=1024,
127
+ depth=8,
128
+ dim_head=64,
129
+ heads=16,
130
+ num_queries=8,
131
+ embedding_dim=768,
132
+ output_dim=1024,
133
+ ff_mult=4,
134
+ ):
135
+ super().__init__()
136
+
137
+ self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5)
138
+
139
+ self.proj_in = nn.Linear(embedding_dim, dim)
140
+
141
+ self.proj_out = nn.Linear(dim, output_dim)
142
+ self.norm_out = nn.LayerNorm(output_dim)
143
+
144
+ self.layers = nn.ModuleList([])
145
+ for _ in range(depth):
146
+ self.layers.append(
147
+ nn.ModuleList(
148
+ [
149
+ PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
150
+ FeedForward(dim=dim, mult=ff_mult),
151
+ ]
152
+ )
153
+ )
154
+
155
+ def forward(self, x):
156
+ latents = self.latents.repeat(x.size(0), 1, 1)
157
+ x = self.proj_in(x)
158
+
159
+ for attn, ff in self.layers:
160
+ latents = attn(x, latents) + latents
161
+ latents = ff(latents) + latents
162
+
163
+ latents = self.proj_out(latents)
164
+ return self.norm_out(latents)
165
+
166
+
167
+ class AttnProcessor(nn.Module):
168
+ r"""
169
+ Default processor for performing attention-related computations.
170
+ """
171
+
172
+ def __init__(
173
+ self,
174
+ hidden_size=None,
175
+ cross_attention_dim=None,
176
+ ):
177
+ super().__init__()
178
+
179
+ def __call__(
180
+ self,
181
+ attn,
182
+ hidden_states,
183
+ encoder_hidden_states=None,
184
+ attention_mask=None,
185
+ temb=None,
186
+ ):
187
+ residual = hidden_states
188
+
189
+ if attn.spatial_norm is not None:
190
+ hidden_states = attn.spatial_norm(hidden_states, temb)
191
+
192
+ input_ndim = hidden_states.ndim
193
+
194
+ if input_ndim == 4:
195
+ batch_size, channel, height, width = hidden_states.shape
196
+ hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
197
+
198
+ batch_size, sequence_length, _ = (
199
+ hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
200
+ )
201
+ attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
202
+
203
+ if attn.group_norm is not None:
204
+ hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
205
+
206
+ query = attn.to_q(hidden_states)
207
+
208
+ if encoder_hidden_states is None:
209
+ encoder_hidden_states = hidden_states
210
+ elif attn.norm_cross:
211
+ encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
212
+
213
+ key = attn.to_k(encoder_hidden_states)
214
+ value = attn.to_v(encoder_hidden_states)
215
+
216
+ query = attn.head_to_batch_dim(query)
217
+ key = attn.head_to_batch_dim(key)
218
+ value = attn.head_to_batch_dim(value)
219
+
220
+ attention_probs = attn.get_attention_scores(query, key, attention_mask)
221
+ hidden_states = torch.bmm(attention_probs, value)
222
+ hidden_states = attn.batch_to_head_dim(hidden_states)
223
+
224
+ # linear proj
225
+ hidden_states = attn.to_out[0](hidden_states)
226
+ # dropout
227
+ hidden_states = attn.to_out[1](hidden_states)
228
+
229
+ if input_ndim == 4:
230
+ hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
231
+
232
+ if attn.residual_connection:
233
+ hidden_states = hidden_states + residual
234
+
235
+ hidden_states = hidden_states / attn.rescale_output_factor
236
+
237
+ return hidden_states
238
+
239
+
240
+ class IPAttnProcessor(nn.Module):
241
+ r"""
242
+ Attention processor for IP-Adapater.
243
+ Args:
244
+ hidden_size (`int`):
245
+ The hidden size of the attention layer.
246
+ cross_attention_dim (`int`):
247
+ The number of channels in the `encoder_hidden_states`.
248
+ scale (`float`, defaults to 1.0):
249
+ the weight scale of image prompt.
250
+ num_tokens (`int`, defaults to 4 when do ip_adapter_plus it should be 16):
251
+ The context length of the image features.
252
+ """
253
+
254
+ def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, num_tokens=4):
255
+ super().__init__()
256
+
257
+ self.hidden_size = hidden_size
258
+ self.cross_attention_dim = cross_attention_dim
259
+ self.scale = scale
260
+ self.num_tokens = num_tokens
261
+
262
+ self.to_k_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
263
+ self.to_v_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
264
+
265
+ def __call__(
266
+ self,
267
+ attn,
268
+ hidden_states,
269
+ encoder_hidden_states=None,
270
+ attention_mask=None,
271
+ temb=None,
272
+ ):
273
+ residual = hidden_states
274
+
275
+ if attn.spatial_norm is not None:
276
+ hidden_states = attn.spatial_norm(hidden_states, temb)
277
+
278
+ input_ndim = hidden_states.ndim
279
+
280
+ if input_ndim == 4:
281
+ batch_size, channel, height, width = hidden_states.shape
282
+ hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
283
+
284
+ batch_size, sequence_length, _ = (
285
+ hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
286
+ )
287
+ attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
288
+
289
+ if attn.group_norm is not None:
290
+ hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
291
+
292
+ query = attn.to_q(hidden_states)
293
+
294
+ if encoder_hidden_states is None:
295
+ encoder_hidden_states = hidden_states
296
+ else:
297
+ # get encoder_hidden_states, ip_hidden_states
298
+ end_pos = encoder_hidden_states.shape[1] - self.num_tokens
299
+ encoder_hidden_states, ip_hidden_states = (
300
+ encoder_hidden_states[:, :end_pos, :],
301
+ encoder_hidden_states[:, end_pos:, :],
302
+ )
303
+ if attn.norm_cross:
304
+ encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
305
+
306
+ key = attn.to_k(encoder_hidden_states)
307
+ value = attn.to_v(encoder_hidden_states)
308
+
309
+ query = attn.head_to_batch_dim(query)
310
+ key = attn.head_to_batch_dim(key)
311
+ value = attn.head_to_batch_dim(value)
312
+
313
+ if xformers_available:
314
+ hidden_states = self._memory_efficient_attention_xformers(query, key, value, attention_mask)
315
+ else:
316
+ attention_probs = attn.get_attention_scores(query, key, attention_mask)
317
+ hidden_states = torch.bmm(attention_probs, value)
318
+ hidden_states = attn.batch_to_head_dim(hidden_states)
319
+
320
+ # for ip-adapter
321
+ ip_key = self.to_k_ip(ip_hidden_states)
322
+ ip_value = self.to_v_ip(ip_hidden_states)
323
+
324
+ ip_key = attn.head_to_batch_dim(ip_key)
325
+ ip_value = attn.head_to_batch_dim(ip_value)
326
+
327
+ if xformers_available:
328
+ ip_hidden_states = self._memory_efficient_attention_xformers(query, ip_key, ip_value, None)
329
+ else:
330
+ ip_attention_probs = attn.get_attention_scores(query, ip_key, None)
331
+ ip_hidden_states = torch.bmm(ip_attention_probs, ip_value)
332
+ ip_hidden_states = attn.batch_to_head_dim(ip_hidden_states)
333
+
334
+ hidden_states = hidden_states + self.scale * ip_hidden_states
335
+
336
+ # linear proj
337
+ hidden_states = attn.to_out[0](hidden_states)
338
+ # dropout
339
+ hidden_states = attn.to_out[1](hidden_states)
340
+
341
+ if input_ndim == 4:
342
+ hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
343
+
344
+ if attn.residual_connection:
345
+ hidden_states = hidden_states + residual
346
+
347
+ hidden_states = hidden_states / attn.rescale_output_factor
348
+
349
+ return hidden_states
350
+
351
+ def _memory_efficient_attention_xformers(self, query, key, value, attention_mask):
352
+ # TODO attention_mask
353
+ query = query.contiguous()
354
+ key = key.contiguous()
355
+ value = value.contiguous()
356
+ hidden_states = xformers.ops.memory_efficient_attention(query, key, value, attn_bias=attention_mask)
357
+ return hidden_states
358
+
359
+
360
+ EXAMPLE_DOC_STRING = """
361
+ Examples:
362
+ ```py
363
+ >>> # !pip install opencv-python transformers accelerate insightface
364
+ >>> import diffusers
365
+ >>> from diffusers.utils import load_image
366
+ >>> from diffusers.models import ControlNetModel
367
+
368
+ >>> import cv2
369
+ >>> import torch
370
+ >>> import numpy as np
371
+ >>> from PIL import Image
372
+
373
+ >>> from insightface.app import FaceAnalysis
374
+ >>> from pipeline_stable_diffusion_xl_instantid import StableDiffusionXLInstantIDPipeline, draw_kps
375
+
376
+ >>> # download 'antelopev2' under ./models
377
+ >>> app = FaceAnalysis(name='antelopev2', root='./', providers=['CUDAExecutionProvider', 'CPUExecutionProvider'])
378
+ >>> app.prepare(ctx_id=0, det_size=(640, 640))
379
+
380
+ >>> # download models under ./checkpoints
381
+ >>> face_adapter = f'./checkpoints/ip-adapter.bin'
382
+ >>> controlnet_path = f'./checkpoints/ControlNetModel'
383
+
384
+ >>> # load IdentityNet
385
+ >>> controlnet = ControlNetModel.from_pretrained(controlnet_path, torch_dtype=torch.float16)
386
+
387
+ >>> pipe = StableDiffusionXLInstantIDPipeline.from_pretrained(
388
+ ... "stabilityai/stable-diffusion-xl-base-1.0", controlnet=controlnet, torch_dtype=torch.float16
389
+ ... )
390
+ >>> pipe.cuda()
391
+
392
+ >>> # load adapter
393
+ >>> pipe.load_ip_adapter_instantid(face_adapter)
394
+
395
+ >>> prompt = "analog film photo of a man. faded film, desaturated, 35mm photo, grainy, vignette, vintage, Kodachrome, Lomography, stained, highly detailed, found footage, masterpiece, best quality"
396
+ >>> negative_prompt = "(lowres, low quality, worst quality:1.2), (text:1.2), watermark, painting, drawing, illustration, glitch, deformed, mutated, cross-eyed, ugly, disfigured (lowres, low quality, worst quality:1.2), (text:1.2), watermark, painting, drawing, illustration, glitch,deformed, mutated, cross-eyed, ugly, disfigured"
397
+
398
+ >>> # load an image
399
+ >>> image = load_image("your-example.jpg")
400
+
401
+ >>> face_info = app.get(cv2.cvtColor(np.array(face_image), cv2.COLOR_RGB2BGR))[-1]
402
+ >>> face_emb = face_info['embedding']
403
+ >>> face_kps = draw_kps(face_image, face_info['kps'])
404
+
405
+ >>> pipe.set_ip_adapter_scale(0.8)
406
+
407
+ >>> # generate image
408
+ >>> image = pipe(
409
+ ... prompt, image_embeds=face_emb, image=face_kps, controlnet_conditioning_scale=0.8
410
+ ... ).images[0]
411
+ ```
412
+ """
413
+
414
+
415
+ def draw_kps(image_pil, kps, color_list=[(255, 0, 0), (0, 255, 0), (0, 0, 255), (255, 255, 0), (255, 0, 255)]):
416
+ stickwidth = 4
417
+ limbSeq = np.array([[0, 2], [1, 2], [3, 2], [4, 2]])
418
+ kps = np.array(kps)
419
+
420
+ w, h = image_pil.size
421
+ out_img = np.zeros([h, w, 3])
422
+
423
+ for i in range(len(limbSeq)):
424
+ index = limbSeq[i]
425
+ color = color_list[index[0]]
426
+
427
+ x = kps[index][:, 0]
428
+ y = kps[index][:, 1]
429
+ length = ((x[0] - x[1]) ** 2 + (y[0] - y[1]) ** 2) ** 0.5
430
+ angle = math.degrees(math.atan2(y[0] - y[1], x[0] - x[1]))
431
+ polygon = cv2.ellipse2Poly(
432
+ (int(np.mean(x)), int(np.mean(y))), (int(length / 2), stickwidth), int(angle), 0, 360, 1
433
+ )
434
+ out_img = cv2.fillConvexPoly(out_img.copy(), polygon, color)
435
+ out_img = (out_img * 0.6).astype(np.uint8)
436
+
437
+ for idx_kp, kp in enumerate(kps):
438
+ color = color_list[idx_kp]
439
+ x, y = kp
440
+ out_img = cv2.circle(out_img.copy(), (int(x), int(y)), 10, color, -1)
441
+
442
+ out_img_pil = PIL.Image.fromarray(out_img.astype(np.uint8))
443
+ return out_img_pil
444
+
445
+
446
+ class StableDiffusionXLInstantIDPipeline(StableDiffusionXLControlNetPipeline):
447
+ def cuda(self, dtype=torch.float16, use_xformers=False):
448
+ self.to("cuda", dtype)
449
+
450
+ if hasattr(self, "image_proj_model"):
451
+ self.image_proj_model.to(self.unet.device).to(self.unet.dtype)
452
+
453
+ if use_xformers:
454
+ if is_xformers_available():
455
+ import xformers
456
+ from packaging import version
457
+
458
+ xformers_version = version.parse(xformers.__version__)
459
+ if xformers_version == version.parse("0.0.16"):
460
+ logger.warning(
461
+ "xFormers 0.0.16 cannot be used for training in some GPUs. If you observe problems during training, please update xFormers to at least 0.0.17. See https://huggingface.co/docs/diffusers/main/en/optimization/xformers for more details."
462
+ )
463
+ self.enable_xformers_memory_efficient_attention()
464
+ else:
465
+ raise ValueError("xformers is not available. Make sure it is installed correctly")
466
+
467
+ def load_ip_adapter_instantid(self, model_ckpt, image_emb_dim=512, num_tokens=16, scale=0.5):
468
+ self.set_image_proj_model(model_ckpt, image_emb_dim, num_tokens)
469
+ self.set_ip_adapter(model_ckpt, num_tokens, scale)
470
+
471
+ def set_image_proj_model(self, model_ckpt, image_emb_dim=512, num_tokens=16):
472
+ image_proj_model = Resampler(
473
+ dim=1280,
474
+ depth=4,
475
+ dim_head=64,
476
+ heads=20,
477
+ num_queries=num_tokens,
478
+ embedding_dim=image_emb_dim,
479
+ output_dim=self.unet.config.cross_attention_dim,
480
+ ff_mult=4,
481
+ )
482
+
483
+ image_proj_model.eval()
484
+
485
+ self.image_proj_model = image_proj_model.to(self.device, dtype=self.dtype)
486
+ state_dict = torch.load(model_ckpt, map_location="cpu")
487
+ if "image_proj" in state_dict:
488
+ state_dict = state_dict["image_proj"]
489
+ self.image_proj_model.load_state_dict(state_dict)
490
+
491
+ self.image_proj_model_in_features = image_emb_dim
492
+
493
+ def set_ip_adapter(self, model_ckpt, num_tokens, scale):
494
+ unet = self.unet
495
+ attn_procs = {}
496
+ for name in unet.attn_processors.keys():
497
+ cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
498
+ if name.startswith("mid_block"):
499
+ hidden_size = unet.config.block_out_channels[-1]
500
+ elif name.startswith("up_blocks"):
501
+ block_id = int(name[len("up_blocks.")])
502
+ hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
503
+ elif name.startswith("down_blocks"):
504
+ block_id = int(name[len("down_blocks.")])
505
+ hidden_size = unet.config.block_out_channels[block_id]
506
+ if cross_attention_dim is None:
507
+ attn_procs[name] = AttnProcessor().to(unet.device, dtype=unet.dtype)
508
+ else:
509
+ attn_procs[name] = IPAttnProcessor(
510
+ hidden_size=hidden_size,
511
+ cross_attention_dim=cross_attention_dim,
512
+ scale=scale,
513
+ num_tokens=num_tokens,
514
+ ).to(unet.device, dtype=unet.dtype)
515
+ unet.set_attn_processor(attn_procs)
516
+
517
+ state_dict = torch.load(model_ckpt, map_location="cpu")
518
+ ip_layers = torch.nn.ModuleList(self.unet.attn_processors.values())
519
+ if "ip_adapter" in state_dict:
520
+ state_dict = state_dict["ip_adapter"]
521
+ ip_layers.load_state_dict(state_dict)
522
+
523
+ def set_ip_adapter_scale(self, scale):
524
+ unet = getattr(self, self.unet_name) if not hasattr(self, "unet") else self.unet
525
+ for attn_processor in unet.attn_processors.values():
526
+ if isinstance(attn_processor, IPAttnProcessor):
527
+ attn_processor.scale = scale
528
+
529
+ def _encode_prompt_image_emb(self, prompt_image_emb, device, dtype, do_classifier_free_guidance):
530
+ if isinstance(prompt_image_emb, torch.Tensor):
531
+ prompt_image_emb = prompt_image_emb.clone().detach()
532
+ else:
533
+ prompt_image_emb = torch.tensor(prompt_image_emb)
534
+
535
+ prompt_image_emb = prompt_image_emb.to(device=device, dtype=dtype)
536
+ prompt_image_emb = prompt_image_emb.reshape([1, -1, self.image_proj_model_in_features])
537
+
538
+ if do_classifier_free_guidance:
539
+ prompt_image_emb = torch.cat([torch.zeros_like(prompt_image_emb), prompt_image_emb], dim=0)
540
+ else:
541
+ prompt_image_emb = torch.cat([prompt_image_emb], dim=0)
542
+
543
+ prompt_image_emb = self.image_proj_model(prompt_image_emb)
544
+ return prompt_image_emb
545
+
546
+ @torch.no_grad()
547
+ @replace_example_docstring(EXAMPLE_DOC_STRING)
548
+ def __call__(
549
+ self,
550
+ prompt: Union[str, List[str]] = None,
551
+ prompt_2: Optional[Union[str, List[str]]] = None,
552
+ image: PipelineImageInput = None,
553
+ height: Optional[int] = None,
554
+ width: Optional[int] = None,
555
+ num_inference_steps: int = 50,
556
+ guidance_scale: float = 5.0,
557
+ negative_prompt: Optional[Union[str, List[str]]] = None,
558
+ negative_prompt_2: Optional[Union[str, List[str]]] = None,
559
+ num_images_per_prompt: Optional[int] = 1,
560
+ eta: float = 0.0,
561
+ generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
562
+ latents: Optional[torch.Tensor] = None,
563
+ prompt_embeds: Optional[torch.Tensor] = None,
564
+ negative_prompt_embeds: Optional[torch.Tensor] = None,
565
+ pooled_prompt_embeds: Optional[torch.Tensor] = None,
566
+ negative_pooled_prompt_embeds: Optional[torch.Tensor] = None,
567
+ image_embeds: Optional[torch.Tensor] = None,
568
+ output_type: str | None = "pil",
569
+ return_dict: bool = True,
570
+ cross_attention_kwargs: Optional[Dict[str, Any]] = None,
571
+ controlnet_conditioning_scale: Union[float, List[float]] = 1.0,
572
+ guess_mode: bool = False,
573
+ control_guidance_start: Union[float, List[float]] = 0.0,
574
+ control_guidance_end: Union[float, List[float]] = 1.0,
575
+ original_size: Tuple[int, int] = None,
576
+ crops_coords_top_left: Tuple[int, int] = (0, 0),
577
+ target_size: Tuple[int, int] = None,
578
+ negative_original_size: Optional[Tuple[int, int]] = None,
579
+ negative_crops_coords_top_left: Tuple[int, int] = (0, 0),
580
+ negative_target_size: Optional[Tuple[int, int]] = None,
581
+ clip_skip: Optional[int] = None,
582
+ callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
583
+ callback_on_step_end_tensor_inputs: List[str] = ["latents"],
584
+ **kwargs,
585
+ ):
586
+ r"""
587
+ The call function to the pipeline for generation.
588
+
589
+ Args:
590
+ prompt (`str` or `List[str]`, *optional*):
591
+ The prompt or prompts to guide image generation. If not defined, you need to pass `prompt_embeds`.
592
+ prompt_2 (`str` or `List[str]`, *optional*):
593
+ The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
594
+ used in both text-encoders.
595
+ image (`torch.Tensor`, `PIL.Image.Image`, `np.ndarray`, `List[torch.Tensor]`, `List[PIL.Image.Image]`, `List[np.ndarray]`,:
596
+ `List[List[torch.Tensor]]`, `List[List[np.ndarray]]` or `List[List[PIL.Image.Image]]`):
597
+ The ControlNet input condition to provide guidance to the `unet` for generation. If the type is
598
+ specified as `torch.Tensor`, it is passed to ControlNet as is. `PIL.Image.Image` can also be
599
+ accepted as an image. The dimensions of the output image defaults to `image`'s dimensions. If height
600
+ and/or width are passed, `image` is resized accordingly. If multiple ControlNets are specified in
601
+ `init`, images must be passed as a list such that each element of the list can be correctly batched for
602
+ input to a single ControlNet.
603
+ height (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
604
+ The height in pixels of the generated image. Anything below 512 pixels won't work well for
605
+ [stabilityai/stable-diffusion-xl-base-1.0](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0)
606
+ and checkpoints that are not specifically fine-tuned on low resolutions.
607
+ width (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
608
+ The width in pixels of the generated image. Anything below 512 pixels won't work well for
609
+ [stabilityai/stable-diffusion-xl-base-1.0](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0)
610
+ and checkpoints that are not specifically fine-tuned on low resolutions.
611
+ num_inference_steps (`int`, *optional*, defaults to 50):
612
+ The number of denoising steps. More denoising steps usually lead to a higher quality image at the
613
+ expense of slower inference.
614
+ guidance_scale (`float`, *optional*, defaults to 5.0):
615
+ A higher guidance scale value encourages the model to generate images closely linked to the text
616
+ `prompt` at the expense of lower image quality. Guidance scale is enabled when `guidance_scale > 1`.
617
+ negative_prompt (`str` or `List[str]`, *optional*):
618
+ The prompt or prompts to guide what to not include in image generation. If not defined, you need to
619
+ pass `negative_prompt_embeds` instead. Ignored when not using guidance (`guidance_scale < 1`).
620
+ negative_prompt_2 (`str` or `List[str]`, *optional*):
621
+ The prompt or prompts to guide what to not include in image generation. This is sent to `tokenizer_2`
622
+ and `text_encoder_2`. If not defined, `negative_prompt` is used in both text-encoders.
623
+ num_images_per_prompt (`int`, *optional*, defaults to 1):
624
+ The number of images to generate per prompt.
625
+ eta (`float`, *optional*, defaults to 0.0):
626
+ Corresponds to parameter eta (η) from the [DDIM](https://huggingface.co/papers/2010.02502) paper. Only applies
627
+ to the [`~schedulers.DDIMScheduler`], and is ignored in other schedulers.
628
+ generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
629
+ A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
630
+ generation deterministic.
631
+ latents (`torch.Tensor`, *optional*):
632
+ Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
633
+ generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
634
+ tensor is generated by sampling using the supplied random `generator`.
635
+ prompt_embeds (`torch.Tensor`, *optional*):
636
+ Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
637
+ provided, text embeddings are generated from the `prompt` input argument.
638
+ negative_prompt_embeds (`torch.Tensor`, *optional*):
639
+ Pre-generated negative text embeddings. Can be used to easily tweak text inputs (prompt weighting). If
640
+ not provided, `negative_prompt_embeds` are generated from the `negative_prompt` input argument.
641
+ pooled_prompt_embeds (`torch.Tensor`, *optional*):
642
+ Pre-generated pooled text embeddings. Can be used to easily tweak text inputs (prompt weighting). If
643
+ not provided, pooled text embeddings are generated from `prompt` input argument.
644
+ negative_pooled_prompt_embeds (`torch.Tensor`, *optional*):
645
+ Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs (prompt
646
+ weighting). If not provided, pooled `negative_prompt_embeds` are generated from `negative_prompt` input
647
+ argument.
648
+ image_embeds (`torch.Tensor`, *optional*):
649
+ Pre-generated image embeddings.
650
+ output_type (`str`, *optional*, defaults to `"pil"`):
651
+ The output format of the generated image. Choose between `PIL.Image` or `np.array`.
652
+ return_dict (`bool`, *optional*, defaults to `True`):
653
+ Whether or not to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a
654
+ plain tuple.
655
+ cross_attention_kwargs (`dict`, *optional*):
656
+ A kwargs dictionary that if specified is passed along to the [`AttentionProcessor`] as defined in
657
+ [`self.processor`](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
658
+ controlnet_conditioning_scale (`float` or `List[float]`, *optional*, defaults to 1.0):
659
+ The outputs of the ControlNet are multiplied by `controlnet_conditioning_scale` before they are added
660
+ to the residual in the original `unet`. If multiple ControlNets are specified in `init`, you can set
661
+ the corresponding scale as a list.
662
+ guess_mode (`bool`, *optional*, defaults to `False`):
663
+ The ControlNet encoder tries to recognize the content of the input image even if you remove all
664
+ prompts. A `guidance_scale` value between 3.0 and 5.0 is recommended.
665
+ control_guidance_start (`float` or `List[float]`, *optional*, defaults to 0.0):
666
+ The percentage of total steps at which the ControlNet starts applying.
667
+ control_guidance_end (`float` or `List[float]`, *optional*, defaults to 1.0):
668
+ The percentage of total steps at which the ControlNet stops applying.
669
+ original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
670
+ If `original_size` is not the same as `target_size` the image will appear to be down- or upsampled.
671
+ `original_size` defaults to `(height, width)` if not specified. Part of SDXL's micro-conditioning as
672
+ explained in section 2.2 of
673
+ [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
674
+ crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)):
675
+ `crops_coords_top_left` can be used to generate an image that appears to be "cropped" from the position
676
+ `crops_coords_top_left` downwards. Favorable, well-centered images are usually achieved by setting
677
+ `crops_coords_top_left` to (0, 0). Part of SDXL's micro-conditioning as explained in section 2.2 of
678
+ [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
679
+ target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
680
+ For most cases, `target_size` should be set to the desired height and width of the generated image. If
681
+ not specified it will default to `(height, width)`. Part of SDXL's micro-conditioning as explained in
682
+ section 2.2 of [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
683
+ negative_original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
684
+ To negatively condition the generation process based on a specific image resolution. Part of SDXL's
685
+ micro-conditioning as explained in section 2.2 of
686
+ [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more
687
+ information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208.
688
+ negative_crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)):
689
+ To negatively condition the generation process based on a specific crop coordinates. Part of SDXL's
690
+ micro-conditioning as explained in section 2.2 of
691
+ [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more
692
+ information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208.
693
+ negative_target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
694
+ To negatively condition the generation process based on a target image resolution. It should be as same
695
+ as the `target_size` for most cases. Part of SDXL's micro-conditioning as explained in section 2.2 of
696
+ [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more
697
+ information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208.
698
+ clip_skip (`int`, *optional*):
699
+ Number of layers to be skipped from CLIP while computing the prompt embeddings. A value of 1 means that
700
+ the output of the pre-final layer will be used for computing the prompt embeddings.
701
+ callback_on_step_end (`Callable`, *optional*):
702
+ A function that calls at the end of each denoising steps during the inference. The function is called
703
+ with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
704
+ callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
705
+ `callback_on_step_end_tensor_inputs`.
706
+ callback_on_step_end_tensor_inputs (`List`, *optional*):
707
+ The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
708
+ will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
709
+ `._callback_tensor_inputs` attribute of your pipeline class.
710
+
711
+ Examples:
712
+
713
+ Returns:
714
+ [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] or `tuple`:
715
+ If `return_dict` is `True`, [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] is returned,
716
+ otherwise a `tuple` is returned containing the output images.
717
+ """
718
+
719
+ callback = kwargs.pop("callback", None)
720
+ callback_steps = kwargs.pop("callback_steps", None)
721
+
722
+ if callback is not None:
723
+ deprecate(
724
+ "callback",
725
+ "1.0.0",
726
+ "Passing `callback` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`",
727
+ )
728
+ if callback_steps is not None:
729
+ deprecate(
730
+ "callback_steps",
731
+ "1.0.0",
732
+ "Passing `callback_steps` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`",
733
+ )
734
+
735
+ controlnet = self.controlnet._orig_mod if is_compiled_module(self.controlnet) else self.controlnet
736
+
737
+ # align format for control guidance
738
+ if not isinstance(control_guidance_start, list) and isinstance(control_guidance_end, list):
739
+ control_guidance_start = len(control_guidance_end) * [control_guidance_start]
740
+ elif not isinstance(control_guidance_end, list) and isinstance(control_guidance_start, list):
741
+ control_guidance_end = len(control_guidance_start) * [control_guidance_end]
742
+ elif not isinstance(control_guidance_start, list) and not isinstance(control_guidance_end, list):
743
+ mult = len(controlnet.nets) if isinstance(controlnet, MultiControlNetModel) else 1
744
+ control_guidance_start, control_guidance_end = (
745
+ mult * [control_guidance_start],
746
+ mult * [control_guidance_end],
747
+ )
748
+
749
+ # 1. Check inputs. Raise error if not correct
750
+ self.check_inputs(
751
+ prompt=prompt,
752
+ prompt_2=prompt_2,
753
+ image=image,
754
+ callback_steps=callback_steps,
755
+ negative_prompt=negative_prompt,
756
+ negative_prompt_2=negative_prompt_2,
757
+ prompt_embeds=prompt_embeds,
758
+ negative_prompt_embeds=negative_prompt_embeds,
759
+ pooled_prompt_embeds=pooled_prompt_embeds,
760
+ negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
761
+ controlnet_conditioning_scale=controlnet_conditioning_scale,
762
+ control_guidance_start=control_guidance_start,
763
+ control_guidance_end=control_guidance_end,
764
+ callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
765
+ )
766
+
767
+ self._guidance_scale = guidance_scale
768
+ self._clip_skip = clip_skip
769
+ self._cross_attention_kwargs = cross_attention_kwargs
770
+
771
+ # 2. Define call parameters
772
+ if prompt is not None and isinstance(prompt, str):
773
+ batch_size = 1
774
+ elif prompt is not None and isinstance(prompt, list):
775
+ batch_size = len(prompt)
776
+ else:
777
+ batch_size = prompt_embeds.shape[0]
778
+
779
+ device = self._execution_device
780
+
781
+ if isinstance(controlnet, MultiControlNetModel) and isinstance(controlnet_conditioning_scale, float):
782
+ controlnet_conditioning_scale = [controlnet_conditioning_scale] * len(controlnet.nets)
783
+
784
+ global_pool_conditions = (
785
+ controlnet.config.global_pool_conditions
786
+ if isinstance(controlnet, ControlNetModel)
787
+ else controlnet.nets[0].config.global_pool_conditions
788
+ )
789
+ guess_mode = guess_mode or global_pool_conditions
790
+
791
+ # 3.1 Encode input prompt
792
+ text_encoder_lora_scale = (
793
+ self.cross_attention_kwargs.get("scale", None) if self.cross_attention_kwargs is not None else None
794
+ )
795
+ (
796
+ prompt_embeds,
797
+ negative_prompt_embeds,
798
+ pooled_prompt_embeds,
799
+ negative_pooled_prompt_embeds,
800
+ ) = self.encode_prompt(
801
+ prompt,
802
+ prompt_2,
803
+ device,
804
+ num_images_per_prompt,
805
+ self.do_classifier_free_guidance,
806
+ negative_prompt,
807
+ negative_prompt_2,
808
+ prompt_embeds=prompt_embeds,
809
+ negative_prompt_embeds=negative_prompt_embeds,
810
+ pooled_prompt_embeds=pooled_prompt_embeds,
811
+ negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
812
+ lora_scale=text_encoder_lora_scale,
813
+ clip_skip=self.clip_skip,
814
+ )
815
+
816
+ # 3.2 Encode image prompt
817
+ prompt_image_emb = self._encode_prompt_image_emb(
818
+ image_embeds, device, self.unet.dtype, self.do_classifier_free_guidance
819
+ )
820
+ bs_embed, seq_len, _ = prompt_image_emb.shape
821
+ prompt_image_emb = prompt_image_emb.repeat(1, num_images_per_prompt, 1)
822
+ prompt_image_emb = prompt_image_emb.view(bs_embed * num_images_per_prompt, seq_len, -1)
823
+
824
+ # 4. Prepare image
825
+ if isinstance(controlnet, ControlNetModel):
826
+ image = self.prepare_image(
827
+ image=image,
828
+ width=width,
829
+ height=height,
830
+ batch_size=batch_size * num_images_per_prompt,
831
+ num_images_per_prompt=num_images_per_prompt,
832
+ device=device,
833
+ dtype=controlnet.dtype,
834
+ do_classifier_free_guidance=self.do_classifier_free_guidance,
835
+ guess_mode=guess_mode,
836
+ )
837
+ height, width = image.shape[-2:]
838
+ elif isinstance(controlnet, MultiControlNetModel):
839
+ images = []
840
+
841
+ for image_ in image:
842
+ image_ = self.prepare_image(
843
+ image=image_,
844
+ width=width,
845
+ height=height,
846
+ batch_size=batch_size * num_images_per_prompt,
847
+ num_images_per_prompt=num_images_per_prompt,
848
+ device=device,
849
+ dtype=controlnet.dtype,
850
+ do_classifier_free_guidance=self.do_classifier_free_guidance,
851
+ guess_mode=guess_mode,
852
+ )
853
+
854
+ images.append(image_)
855
+
856
+ image = images
857
+ height, width = image[0].shape[-2:]
858
+ else:
859
+ assert False
860
+
861
+ # 5. Prepare timesteps
862
+ self.scheduler.set_timesteps(num_inference_steps, device=device)
863
+ timesteps = self.scheduler.timesteps
864
+ self._num_timesteps = len(timesteps)
865
+
866
+ # 6. Prepare latent variables
867
+ num_channels_latents = self.unet.config.in_channels
868
+ latents = self.prepare_latents(
869
+ batch_size * num_images_per_prompt,
870
+ num_channels_latents,
871
+ height,
872
+ width,
873
+ prompt_embeds.dtype,
874
+ device,
875
+ generator,
876
+ latents,
877
+ )
878
+
879
+ # 6.5 Optionally get Guidance Scale Embedding
880
+ timestep_cond = None
881
+ if self.unet.config.time_cond_proj_dim is not None:
882
+ guidance_scale_tensor = torch.tensor(self.guidance_scale - 1).repeat(batch_size * num_images_per_prompt)
883
+ timestep_cond = self.get_guidance_scale_embedding(
884
+ guidance_scale_tensor, embedding_dim=self.unet.config.time_cond_proj_dim
885
+ ).to(device=device, dtype=latents.dtype)
886
+
887
+ # 7. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
888
+ extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
889
+
890
+ # 7.1 Create tensor stating which controlnets to keep
891
+ controlnet_keep = []
892
+ for i in range(len(timesteps)):
893
+ keeps = [
894
+ 1.0 - float(i / len(timesteps) < s or (i + 1) / len(timesteps) > e)
895
+ for s, e in zip(control_guidance_start, control_guidance_end)
896
+ ]
897
+ controlnet_keep.append(keeps[0] if isinstance(controlnet, ControlNetModel) else keeps)
898
+
899
+ # 7.2 Prepare added time ids & embeddings
900
+ if isinstance(image, list):
901
+ original_size = original_size or image[0].shape[-2:]
902
+ else:
903
+ original_size = original_size or image.shape[-2:]
904
+ target_size = target_size or (height, width)
905
+
906
+ add_text_embeds = pooled_prompt_embeds
907
+ if self.text_encoder_2 is None:
908
+ text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
909
+ else:
910
+ text_encoder_projection_dim = self.text_encoder_2.config.projection_dim
911
+
912
+ add_time_ids = self._get_add_time_ids(
913
+ original_size,
914
+ crops_coords_top_left,
915
+ target_size,
916
+ dtype=prompt_embeds.dtype,
917
+ text_encoder_projection_dim=text_encoder_projection_dim,
918
+ )
919
+
920
+ if negative_original_size is not None and negative_target_size is not None:
921
+ negative_add_time_ids = self._get_add_time_ids(
922
+ negative_original_size,
923
+ negative_crops_coords_top_left,
924
+ negative_target_size,
925
+ dtype=prompt_embeds.dtype,
926
+ text_encoder_projection_dim=text_encoder_projection_dim,
927
+ )
928
+ else:
929
+ negative_add_time_ids = add_time_ids
930
+
931
+ if self.do_classifier_free_guidance:
932
+ prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
933
+ add_text_embeds = torch.cat([negative_pooled_prompt_embeds, add_text_embeds], dim=0)
934
+ add_time_ids = torch.cat([negative_add_time_ids, add_time_ids], dim=0)
935
+
936
+ prompt_embeds = prompt_embeds.to(device)
937
+ add_text_embeds = add_text_embeds.to(device)
938
+ add_time_ids = add_time_ids.to(device).repeat(batch_size * num_images_per_prompt, 1)
939
+ encoder_hidden_states = torch.cat([prompt_embeds, prompt_image_emb], dim=1)
940
+
941
+ # 8. Denoising loop
942
+ num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
943
+ is_unet_compiled = is_compiled_module(self.unet)
944
+ is_controlnet_compiled = is_compiled_module(self.controlnet)
945
+ is_torch_higher_equal_2_1 = is_torch_version(">=", "2.1")
946
+
947
+ with self.progress_bar(total=num_inference_steps) as progress_bar:
948
+ for i, t in enumerate(timesteps):
949
+ # Relevant thread:
950
+ # https://dev-discuss.pytorch.org/t/cudagraphs-in-pytorch-2-0/1428
951
+ if (is_unet_compiled and is_controlnet_compiled) and is_torch_higher_equal_2_1:
952
+ torch._inductor.cudagraph_mark_step_begin()
953
+ # expand the latents if we are doing classifier free guidance
954
+ latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
955
+ latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
956
+
957
+ added_cond_kwargs = {"text_embeds": add_text_embeds, "time_ids": add_time_ids}
958
+
959
+ # controlnet(s) inference
960
+ if guess_mode and self.do_classifier_free_guidance:
961
+ # Infer ControlNet only for the conditional batch.
962
+ control_model_input = latents
963
+ control_model_input = self.scheduler.scale_model_input(control_model_input, t)
964
+ controlnet_prompt_embeds = prompt_embeds.chunk(2)[1]
965
+ controlnet_added_cond_kwargs = {
966
+ "text_embeds": add_text_embeds.chunk(2)[1],
967
+ "time_ids": add_time_ids.chunk(2)[1],
968
+ }
969
+ else:
970
+ control_model_input = latent_model_input
971
+ controlnet_prompt_embeds = prompt_embeds
972
+ controlnet_added_cond_kwargs = added_cond_kwargs
973
+
974
+ if isinstance(controlnet_keep[i], list):
975
+ cond_scale = [c * s for c, s in zip(controlnet_conditioning_scale, controlnet_keep[i])]
976
+ else:
977
+ controlnet_cond_scale = controlnet_conditioning_scale
978
+ if isinstance(controlnet_cond_scale, list):
979
+ controlnet_cond_scale = controlnet_cond_scale[0]
980
+ cond_scale = controlnet_cond_scale * controlnet_keep[i]
981
+
982
+ down_block_res_samples, mid_block_res_sample = self.controlnet(
983
+ control_model_input,
984
+ t,
985
+ encoder_hidden_states=prompt_image_emb,
986
+ controlnet_cond=image,
987
+ conditioning_scale=cond_scale,
988
+ guess_mode=guess_mode,
989
+ added_cond_kwargs=controlnet_added_cond_kwargs,
990
+ return_dict=False,
991
+ )
992
+
993
+ if guess_mode and self.do_classifier_free_guidance:
994
+ # Inferred ControlNet only for the conditional batch.
995
+ # To apply the output of ControlNet to both the unconditional and conditional batches,
996
+ # add 0 to the unconditional batch to keep it unchanged.
997
+ down_block_res_samples = [torch.cat([torch.zeros_like(d), d]) for d in down_block_res_samples]
998
+ mid_block_res_sample = torch.cat([torch.zeros_like(mid_block_res_sample), mid_block_res_sample])
999
+
1000
+ # predict the noise residual
1001
+ noise_pred = self.unet(
1002
+ latent_model_input,
1003
+ t,
1004
+ encoder_hidden_states=encoder_hidden_states,
1005
+ timestep_cond=timestep_cond,
1006
+ cross_attention_kwargs=self.cross_attention_kwargs,
1007
+ down_block_additional_residuals=down_block_res_samples,
1008
+ mid_block_additional_residual=mid_block_res_sample,
1009
+ added_cond_kwargs=added_cond_kwargs,
1010
+ return_dict=False,
1011
+ )[0]
1012
+
1013
+ # perform guidance
1014
+ if self.do_classifier_free_guidance:
1015
+ noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
1016
+ noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
1017
+
1018
+ # compute the previous noisy sample x_t -> x_t-1
1019
+ latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
1020
+
1021
+ if callback_on_step_end is not None:
1022
+ callback_kwargs = {}
1023
+ for k in callback_on_step_end_tensor_inputs:
1024
+ callback_kwargs[k] = locals()[k]
1025
+ callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
1026
+
1027
+ latents = callback_outputs.pop("latents", latents)
1028
+ prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
1029
+ negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
1030
+
1031
+ # call the callback, if provided
1032
+ if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
1033
+ progress_bar.update()
1034
+ if callback is not None and i % callback_steps == 0:
1035
+ step_idx = i // getattr(self.scheduler, "order", 1)
1036
+ callback(step_idx, t, latents)
1037
+
1038
+ if not output_type == "latent":
1039
+ # make sure the VAE is in float32 mode, as it overflows in float16
1040
+ needs_upcasting = self.vae.dtype == torch.float16 and self.vae.config.force_upcast
1041
+ if needs_upcasting:
1042
+ self.upcast_vae()
1043
+ latents = latents.to(next(iter(self.vae.post_quant_conv.parameters())).dtype)
1044
+
1045
+ image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0]
1046
+
1047
+ # cast back to fp16 if needed
1048
+ if needs_upcasting:
1049
+ self.vae.to(dtype=torch.float16)
1050
+ else:
1051
+ image = latents
1052
+
1053
+ if not output_type == "latent":
1054
+ # apply watermark if available
1055
+ if self.watermark is not None:
1056
+ image = self.watermark.apply_watermark(image)
1057
+
1058
+ image = self.image_processor.postprocess(image, output_type=output_type)
1059
+
1060
+ # Offload all models
1061
+ self.maybe_free_model_hooks()
1062
+
1063
+ if not return_dict:
1064
+ return (image,)
1065
+
1066
+ return StableDiffusionXLPipelineOutput(images=image)