comdoleger commited on
Commit
c98db2a
·
verified ·
1 Parent(s): 3c8c25c

Upload extensions_built_in/diffusion_models/qwen_image/qwen_image_pipelines.py with huggingface_hub

Browse files
extensions_built_in/diffusion_models/qwen_image/qwen_image_pipelines.py ADDED
@@ -0,0 +1,354 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Any, Callable, Dict, List, Optional, Union
2
+
3
+ import numpy as np
4
+ import torch
5
+
6
+ try:
7
+ from diffusers import QwenImageEditPlusPipeline
8
+ from diffusers.pipelines.qwenimage.pipeline_qwenimage_edit_plus import (
9
+ CONDITION_IMAGE_SIZE,
10
+ VAE_IMAGE_SIZE,
11
+ XLA_AVAILABLE,
12
+ logger,
13
+ calculate_dimensions,
14
+ calculate_shift,
15
+ retrieve_timesteps,
16
+ )
17
+ except ImportError:
18
+ raise ImportError(
19
+ "Diffusers is out of date. Update diffusers to the latest version by doing 'pip uninstall diffusers' and then 'pip install -r requirements.txt'"
20
+ )
21
+
22
+ from diffusers.image_processor import PipelineImageInput
23
+ from diffusers.pipelines.qwenimage.pipeline_output import QwenImagePipelineOutput
24
+
25
+
26
+ class QwenImageEditPlusCustomPipeline(QwenImageEditPlusPipeline):
27
+ @torch.no_grad()
28
+ def __call__(
29
+ self,
30
+ image: Optional[PipelineImageInput] = None,
31
+ prompt: Union[str, List[str]] = None,
32
+ negative_prompt: Union[str, List[str]] = None,
33
+ true_cfg_scale: float = 4.0,
34
+ height: Optional[int] = None,
35
+ width: Optional[int] = None,
36
+ num_inference_steps: int = 50,
37
+ sigmas: Optional[List[float]] = None,
38
+ guidance_scale: Optional[float] = None,
39
+ num_images_per_prompt: int = 1,
40
+ generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
41
+ latents: Optional[torch.Tensor] = None,
42
+ prompt_embeds: Optional[torch.Tensor] = None,
43
+ prompt_embeds_mask: Optional[torch.Tensor] = None,
44
+ negative_prompt_embeds: Optional[torch.Tensor] = None,
45
+ negative_prompt_embeds_mask: Optional[torch.Tensor] = None,
46
+ output_type: Optional[str] = "pil",
47
+ return_dict: bool = True,
48
+ attention_kwargs: Optional[Dict[str, Any]] = None,
49
+ callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
50
+ callback_on_step_end_tensor_inputs: List[str] = ["latents"],
51
+ max_sequence_length: int = 512,
52
+ do_cfg_norm: bool = False,
53
+ ):
54
+ image_size = image[-1].size if isinstance(image, list) else image.size
55
+ calculated_width, calculated_height = calculate_dimensions(
56
+ 1024 * 1024, image_size[0] / image_size[1]
57
+ )
58
+ height = height or calculated_height
59
+ width = width or calculated_width
60
+
61
+ multiple_of = self.vae_scale_factor * 2
62
+ width = width // multiple_of * multiple_of
63
+ height = height // multiple_of * multiple_of
64
+
65
+ # 1. Check inputs. Raise error if not correct
66
+ self.check_inputs(
67
+ prompt,
68
+ height,
69
+ width,
70
+ negative_prompt=negative_prompt,
71
+ prompt_embeds=prompt_embeds,
72
+ negative_prompt_embeds=negative_prompt_embeds,
73
+ prompt_embeds_mask=prompt_embeds_mask,
74
+ negative_prompt_embeds_mask=negative_prompt_embeds_mask,
75
+ callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
76
+ max_sequence_length=max_sequence_length,
77
+ )
78
+
79
+ self._guidance_scale = guidance_scale
80
+ self._attention_kwargs = attention_kwargs
81
+ self._current_timestep = None
82
+ self._interrupt = False
83
+
84
+ # 2. Define call parameters
85
+ if prompt is not None and isinstance(prompt, str):
86
+ batch_size = 1
87
+ elif prompt is not None and isinstance(prompt, list):
88
+ batch_size = len(prompt)
89
+ else:
90
+ batch_size = prompt_embeds.shape[0]
91
+
92
+ device = self._execution_device
93
+ # 3. Preprocess image
94
+ if image is not None and not (
95
+ isinstance(image, torch.Tensor) and image.size(1) == self.latent_channels
96
+ ):
97
+ if not isinstance(image, list):
98
+ image = [image]
99
+ condition_image_sizes = []
100
+ condition_images = []
101
+ vae_image_sizes = []
102
+ vae_images = []
103
+ for img in image:
104
+ image_width, image_height = img.size
105
+ condition_width, condition_height = calculate_dimensions(
106
+ CONDITION_IMAGE_SIZE, image_width / image_height
107
+ )
108
+ vae_width, vae_height = calculate_dimensions(
109
+ VAE_IMAGE_SIZE, image_width / image_height
110
+ )
111
+ condition_image_sizes.append((condition_width, condition_height))
112
+ vae_image_sizes.append((vae_width, vae_height))
113
+ condition_images.append(
114
+ self.image_processor.resize(img, condition_height, condition_width)
115
+ )
116
+ vae_images.append(
117
+ self.image_processor.preprocess(
118
+ img, vae_height, vae_width
119
+ ).unsqueeze(2)
120
+ )
121
+
122
+ has_neg_prompt = negative_prompt is not None or (
123
+ negative_prompt_embeds is not None
124
+ and negative_prompt_embeds_mask is not None
125
+ )
126
+
127
+ if true_cfg_scale > 1 and not has_neg_prompt:
128
+ logger.warning(
129
+ f"true_cfg_scale is passed as {true_cfg_scale}, but classifier-free guidance is not enabled since no negative_prompt is provided."
130
+ )
131
+ elif true_cfg_scale <= 1 and has_neg_prompt:
132
+ logger.warning(
133
+ " negative_prompt is passed but classifier-free guidance is not enabled since true_cfg_scale <= 1"
134
+ )
135
+
136
+ do_true_cfg = true_cfg_scale > 1 and has_neg_prompt
137
+ prompt_embeds, prompt_embeds_mask = self.encode_prompt(
138
+ image=condition_images,
139
+ prompt=prompt,
140
+ prompt_embeds=prompt_embeds,
141
+ prompt_embeds_mask=prompt_embeds_mask,
142
+ device=device,
143
+ num_images_per_prompt=num_images_per_prompt,
144
+ max_sequence_length=max_sequence_length,
145
+ )
146
+ if do_true_cfg:
147
+ negative_prompt_embeds, negative_prompt_embeds_mask = self.encode_prompt(
148
+ image=condition_images,
149
+ prompt=negative_prompt,
150
+ prompt_embeds=negative_prompt_embeds,
151
+ prompt_embeds_mask=negative_prompt_embeds_mask,
152
+ device=device,
153
+ num_images_per_prompt=num_images_per_prompt,
154
+ max_sequence_length=max_sequence_length,
155
+ )
156
+
157
+ # 4. Prepare latent variables
158
+ num_channels_latents = self.transformer.config.in_channels // 4
159
+ latents, image_latents = self.prepare_latents(
160
+ vae_images,
161
+ batch_size * num_images_per_prompt,
162
+ num_channels_latents,
163
+ height,
164
+ width,
165
+ prompt_embeds.dtype,
166
+ device,
167
+ generator,
168
+ latents,
169
+ )
170
+ img_shapes = [
171
+ [
172
+ (
173
+ 1,
174
+ height // self.vae_scale_factor // 2,
175
+ width // self.vae_scale_factor // 2,
176
+ ),
177
+ *[
178
+ (
179
+ 1,
180
+ vae_height // self.vae_scale_factor // 2,
181
+ vae_width // self.vae_scale_factor // 2,
182
+ )
183
+ for vae_width, vae_height in vae_image_sizes
184
+ ],
185
+ ]
186
+ ] * batch_size
187
+
188
+ # 5. Prepare timesteps
189
+ sigmas = (
190
+ np.linspace(1.0, 1 / num_inference_steps, num_inference_steps)
191
+ if sigmas is None
192
+ else sigmas
193
+ )
194
+ image_seq_len = latents.shape[1]
195
+ mu = calculate_shift(
196
+ image_seq_len,
197
+ self.scheduler.config.get("base_image_seq_len", 256),
198
+ self.scheduler.config.get("max_image_seq_len", 4096),
199
+ self.scheduler.config.get("base_shift", 0.5),
200
+ self.scheduler.config.get("max_shift", 1.15),
201
+ )
202
+ timesteps, num_inference_steps = retrieve_timesteps(
203
+ self.scheduler,
204
+ num_inference_steps,
205
+ device,
206
+ sigmas=sigmas,
207
+ mu=mu,
208
+ )
209
+ num_warmup_steps = max(
210
+ len(timesteps) - num_inference_steps * self.scheduler.order, 0
211
+ )
212
+ self._num_timesteps = len(timesteps)
213
+
214
+ # handle guidance
215
+ if self.transformer.config.guidance_embeds and guidance_scale is None:
216
+ raise ValueError("guidance_scale is required for guidance-distilled model.")
217
+ elif self.transformer.config.guidance_embeds:
218
+ guidance = torch.full(
219
+ [1], guidance_scale, device=device, dtype=torch.float32
220
+ )
221
+ guidance = guidance.expand(latents.shape[0])
222
+ elif not self.transformer.config.guidance_embeds and guidance_scale is not None:
223
+ logger.warning(
224
+ f"guidance_scale is passed as {guidance_scale}, but ignored since the model is not guidance-distilled."
225
+ )
226
+ guidance = None
227
+ elif not self.transformer.config.guidance_embeds and guidance_scale is None:
228
+ guidance = None
229
+
230
+ if self.attention_kwargs is None:
231
+ self._attention_kwargs = {}
232
+
233
+ txt_seq_lens = (
234
+ prompt_embeds_mask.sum(dim=1).tolist()
235
+ if prompt_embeds_mask is not None
236
+ else None
237
+ )
238
+ negative_txt_seq_lens = (
239
+ negative_prompt_embeds_mask.sum(dim=1).tolist()
240
+ if negative_prompt_embeds_mask is not None
241
+ else None
242
+ )
243
+
244
+ # 6. Denoising loop
245
+ self.scheduler.set_begin_index(0)
246
+ with self.progress_bar(total=num_inference_steps) as progress_bar:
247
+ for i, t in enumerate(timesteps):
248
+ if self.interrupt:
249
+ continue
250
+
251
+ self._current_timestep = t
252
+
253
+ latent_model_input = latents
254
+ if image_latents is not None:
255
+ latent_model_input = torch.cat([latents, image_latents], dim=1)
256
+
257
+ # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
258
+ timestep = t.expand(latents.shape[0]).to(latents.dtype)
259
+ with self.transformer.cache_context("cond"):
260
+ noise_pred = self.transformer(
261
+ hidden_states=latent_model_input,
262
+ timestep=timestep / 1000,
263
+ guidance=guidance,
264
+ encoder_hidden_states_mask=prompt_embeds_mask,
265
+ encoder_hidden_states=prompt_embeds,
266
+ img_shapes=img_shapes,
267
+ txt_seq_lens=txt_seq_lens,
268
+ attention_kwargs=self.attention_kwargs,
269
+ return_dict=False,
270
+ )[0]
271
+ noise_pred = noise_pred[:, : latents.size(1)]
272
+
273
+ if do_true_cfg:
274
+ with self.transformer.cache_context("uncond"):
275
+ neg_noise_pred = self.transformer(
276
+ hidden_states=latent_model_input,
277
+ timestep=timestep / 1000,
278
+ guidance=guidance,
279
+ encoder_hidden_states_mask=negative_prompt_embeds_mask,
280
+ encoder_hidden_states=negative_prompt_embeds,
281
+ img_shapes=img_shapes,
282
+ txt_seq_lens=negative_txt_seq_lens,
283
+ attention_kwargs=self.attention_kwargs,
284
+ return_dict=False,
285
+ )[0]
286
+ neg_noise_pred = neg_noise_pred[:, : latents.size(1)]
287
+ comb_pred = neg_noise_pred + true_cfg_scale * (
288
+ noise_pred - neg_noise_pred
289
+ )
290
+
291
+ if do_cfg_norm:
292
+ # the official code does this, but I find it hurts more often than it helps, leaving it optional but off by default
293
+ cond_norm = torch.norm(noise_pred, dim=-1, keepdim=True)
294
+ noise_norm = torch.norm(comb_pred, dim=-1, keepdim=True)
295
+ noise_pred = comb_pred * (cond_norm / noise_norm)
296
+ else:
297
+ noise_pred = comb_pred
298
+
299
+ # compute the previous noisy sample x_t -> x_t-1
300
+ latents_dtype = latents.dtype
301
+ latents = self.scheduler.step(
302
+ noise_pred, t, latents, return_dict=False
303
+ )[0]
304
+
305
+ if latents.dtype != latents_dtype:
306
+ if torch.backends.mps.is_available():
307
+ # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
308
+ latents = latents.to(latents_dtype)
309
+
310
+ if callback_on_step_end is not None:
311
+ callback_kwargs = {}
312
+ for k in callback_on_step_end_tensor_inputs:
313
+ callback_kwargs[k] = locals()[k]
314
+ callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
315
+
316
+ latents = callback_outputs.pop("latents", latents)
317
+ prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
318
+
319
+ # call the callback, if provided
320
+ if i == len(timesteps) - 1 or (
321
+ (i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0
322
+ ):
323
+ progress_bar.update()
324
+
325
+ if XLA_AVAILABLE:
326
+ xm.mark_step()
327
+
328
+ self._current_timestep = None
329
+ if output_type == "latent":
330
+ image = latents
331
+ else:
332
+ latents = self._unpack_latents(
333
+ latents, height, width, self.vae_scale_factor
334
+ )
335
+ latents = latents.to(self.vae.dtype)
336
+ latents_mean = (
337
+ torch.tensor(self.vae.config.latents_mean)
338
+ .view(1, self.vae.config.z_dim, 1, 1, 1)
339
+ .to(latents.device, latents.dtype)
340
+ )
341
+ latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(
342
+ 1, self.vae.config.z_dim, 1, 1, 1
343
+ ).to(latents.device, latents.dtype)
344
+ latents = latents / latents_std + latents_mean
345
+ image = self.vae.decode(latents, return_dict=False)[0][:, :, 0]
346
+ image = self.image_processor.postprocess(image, output_type=output_type)
347
+
348
+ # Offload all models
349
+ self.maybe_free_model_hooks()
350
+
351
+ if not return_dict:
352
+ return (image,)
353
+
354
+ return QwenImagePipelineOutput(images=image)