comdoleger commited on
Commit
c832837
·
verified ·
1 Parent(s): ae9d5a5

Upload extensions_built_in/diffusion_models/chroma/pipeline.py with huggingface_hub

Browse files
extensions_built_in/diffusion_models/chroma/pipeline.py ADDED
@@ -0,0 +1,328 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Union, List, Optional, Dict, Any, Callable
2
+
3
+ import numpy as np
4
+ import torch
5
+ from diffusers import FluxPipeline
6
+ from diffusers.pipelines.flux.pipeline_flux import calculate_shift, retrieve_timesteps
7
+ from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput
8
+ from diffusers.utils import is_torch_xla_available
9
+ from diffusers.utils.torch_utils import randn_tensor
10
+
11
+
12
+ if is_torch_xla_available():
13
+ import torch_xla.core.xla_model as xm
14
+
15
+ XLA_AVAILABLE = True
16
+ else:
17
+ XLA_AVAILABLE = False
18
+
19
+
20
+ def prepare_latent_image_ids(batch_size, height, width, patch_size=2, max_offset=0):
21
+ """
22
+ Generates positional embeddings for a latent image.
23
+
24
+ Args:
25
+ batch_size (int): The number of images in the batch.
26
+ height (int): The height of the image.
27
+ width (int): The width of the image.
28
+ patch_size (int, optional): The size of the patches. Defaults to 2.
29
+ max_offset (int, optional): The maximum random offset to apply. Defaults to 0.
30
+
31
+ Returns:
32
+ torch.Tensor: A tensor containing the positional embeddings.
33
+ """
34
+ # the random pos embedding helps generalize to larger res without training at large res
35
+ # pos embedding for rope, 2d pos embedding, corner embedding and not center based
36
+ latent_image_ids = torch.zeros(height // patch_size, width // patch_size, 3)
37
+
38
+ # Add positional encodings
39
+ latent_image_ids[..., 1] = (
40
+ latent_image_ids[..., 1] + torch.arange(height // patch_size)[:, None]
41
+ )
42
+ latent_image_ids[..., 2] = (
43
+ latent_image_ids[..., 2] + torch.arange(width // patch_size)[None, :]
44
+ )
45
+
46
+ # Add random offset if specified
47
+ if max_offset > 0:
48
+ offset_y = torch.randint(0, max_offset + 1, (1,)).item()
49
+ offset_x = torch.randint(0, max_offset + 1, (1,)).item()
50
+ latent_image_ids[..., 1] += offset_y
51
+ latent_image_ids[..., 2] += offset_x
52
+
53
+
54
+ (
55
+ latent_image_id_height,
56
+ latent_image_id_width,
57
+ latent_image_id_channels,
58
+ ) = latent_image_ids.shape
59
+
60
+ # Reshape for batch
61
+ latent_image_ids = latent_image_ids[None, :].repeat(batch_size, 1, 1, 1)
62
+ latent_image_ids = latent_image_ids.reshape(
63
+ batch_size,
64
+ latent_image_id_height * latent_image_id_width,
65
+ latent_image_id_channels,
66
+ )
67
+
68
+ return latent_image_ids
69
+
70
+
71
+ class ChromaPipeline(FluxPipeline):
72
+ def __init__(
73
+ self,
74
+ scheduler,
75
+ vae,
76
+ text_encoder,
77
+ tokenizer,
78
+ text_encoder_2,
79
+ tokenizer_2,
80
+ transformer,
81
+ image_encoder = None,
82
+ feature_extractor = None,
83
+ is_radiance: bool = False,
84
+ ):
85
+ super().__init__(
86
+ scheduler=scheduler,
87
+ vae=vae,
88
+ text_encoder=text_encoder,
89
+ tokenizer=tokenizer,
90
+ text_encoder_2=text_encoder_2,
91
+ tokenizer_2=tokenizer_2,
92
+ transformer=transformer,
93
+ image_encoder=image_encoder,
94
+ feature_extractor=feature_extractor,
95
+ )
96
+ self.is_radiance = is_radiance
97
+ self.vae_scale_factor = 8 if not is_radiance else 1
98
+
99
+ def prepare_latents(
100
+ self,
101
+ batch_size,
102
+ num_channels_latents,
103
+ height,
104
+ width,
105
+ dtype,
106
+ device,
107
+ generator,
108
+ latents=None,
109
+ ):
110
+ # VAE applies 8x compression on images but we must also account for packing which requires
111
+ # latent height and width to be divisible by 2.
112
+ height = 2 * (int(height) // (self.vae_scale_factor * 2))
113
+ width = 2 * (int(width) // (self.vae_scale_factor * 2))
114
+
115
+ shape = (batch_size, num_channels_latents, height, width)
116
+
117
+ if latents is not None:
118
+ latent_image_ids = prepare_latent_image_ids(
119
+ batch_size,
120
+ height,
121
+ width,
122
+ patch_size=2 if not self.is_radiance else 16
123
+ ).to(device=device, dtype=dtype)
124
+ # latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype)
125
+ return latents.to(device=device, dtype=dtype), latent_image_ids
126
+
127
+ if isinstance(generator, list) and len(generator) != batch_size:
128
+ raise ValueError(
129
+ f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
130
+ f" size of {batch_size}. Make sure the batch size matches the length of the generators."
131
+ )
132
+
133
+ latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
134
+
135
+ if not self.is_radiance:
136
+ latents = self._pack_latents(latents, batch_size, num_channels_latents, height, width)
137
+
138
+ # latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype)
139
+ latent_image_ids = prepare_latent_image_ids(
140
+ batch_size,
141
+ height,
142
+ width,
143
+ patch_size=2 if not self.is_radiance else 16
144
+ ).to(device=device, dtype=dtype)
145
+
146
+ return latents, latent_image_ids
147
+
148
+ def __call__(
149
+ self,
150
+ prompt: Union[str, List[str]] = None,
151
+ prompt_2: Optional[Union[str, List[str]]] = None,
152
+ negative_prompt: Optional[Union[str, List[str]]] = None,
153
+ negative_prompt_2: Optional[Union[str, List[str]]] = None,
154
+ height: Optional[int] = None,
155
+ width: Optional[int] = None,
156
+ num_inference_steps: int = 28,
157
+ timesteps: List[int] = None,
158
+ guidance_scale: float = 7.0,
159
+ num_images_per_prompt: Optional[int] = 1,
160
+ generator: Optional[Union[torch.Generator,
161
+ List[torch.Generator]]] = None,
162
+ latents: Optional[torch.FloatTensor] = None,
163
+ prompt_embeds: Optional[torch.FloatTensor] = None,
164
+ prompt_attn_mask: Optional[torch.FloatTensor] = None,
165
+ negative_prompt_embeds: Optional[torch.FloatTensor] = None,
166
+ negative_prompt_attn_mask: Optional[torch.FloatTensor] = None,
167
+ output_type: Optional[str] = "pil",
168
+ return_dict: bool = True,
169
+ joint_attention_kwargs: Optional[Dict[str, Any]] = None,
170
+ callback_on_step_end: Optional[Callable[[
171
+ int, int, Dict], None]] = None,
172
+ callback_on_step_end_tensor_inputs: List[str] = ["latents"],
173
+ max_sequence_length: int = 512,
174
+ ):
175
+
176
+ height = height or self.default_sample_size * self.vae_scale_factor
177
+ width = width or self.default_sample_size * self.vae_scale_factor
178
+
179
+ self._guidance_scale = guidance_scale
180
+ self._joint_attention_kwargs = joint_attention_kwargs
181
+ self._interrupt = False
182
+
183
+ # 2. Define call parameters
184
+ if prompt is not None and isinstance(prompt, str):
185
+ batch_size = 1
186
+ elif prompt is not None and isinstance(prompt, list):
187
+ batch_size = len(prompt)
188
+ else:
189
+ batch_size = prompt_embeds.shape[0]
190
+
191
+ device = self._execution_device
192
+ if isinstance(device, str):
193
+ device = torch.device(device)
194
+
195
+ text_ids = torch.zeros(batch_size, prompt_embeds.shape[1], 3).to(device=device, dtype=torch.bfloat16)
196
+ if guidance_scale > 1.00001:
197
+ negative_text_ids = torch.zeros(batch_size, negative_prompt_embeds.shape[1], 3).to(device=device, dtype=torch.bfloat16)
198
+
199
+ # 4. Prepare latent variables
200
+ num_channels_latents = 64 // 4
201
+ if self.is_radiance:
202
+ num_channels_latents = 3
203
+ latents, latent_image_ids = self.prepare_latents(
204
+ batch_size * num_images_per_prompt,
205
+ num_channels_latents,
206
+ height,
207
+ width,
208
+ prompt_embeds.dtype,
209
+ device,
210
+ generator,
211
+ latents,
212
+ )
213
+
214
+ # extend img ids to match batch size
215
+ # latent_image_ids = latent_image_ids.unsqueeze(0)
216
+ # latent_image_ids = torch.cat([latent_image_ids] * batch_size, dim=0)
217
+
218
+ # 5. Prepare timesteps
219
+ sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps)
220
+ image_seq_len = latents.shape[1]
221
+ mu = calculate_shift(
222
+ image_seq_len,
223
+ self.scheduler.config.base_image_seq_len,
224
+ self.scheduler.config.max_image_seq_len,
225
+ self.scheduler.config.base_shift,
226
+ self.scheduler.config.max_shift,
227
+ )
228
+ timesteps, num_inference_steps = retrieve_timesteps(
229
+ self.scheduler,
230
+ num_inference_steps,
231
+ device,
232
+ timesteps,
233
+ sigmas,
234
+ mu=mu,
235
+ )
236
+ num_warmup_steps = max(
237
+ len(timesteps) - num_inference_steps * self.scheduler.order, 0)
238
+ self._num_timesteps = len(timesteps)
239
+
240
+ guidance = torch.full([1], 0, device=device, dtype=torch.float32)
241
+ guidance = guidance.expand(latents.shape[0])
242
+
243
+ # 6. Denoising loop
244
+ with self.progress_bar(total=num_inference_steps) as progress_bar:
245
+ for i, t in enumerate(timesteps):
246
+ if self.interrupt:
247
+ continue
248
+
249
+ # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
250
+ timestep = t.expand(latents.shape[0]).to(latents.dtype)
251
+
252
+ # handle guidance
253
+
254
+ noise_pred_text = self.transformer(
255
+ img=latents,
256
+ img_ids=latent_image_ids,
257
+ txt=prompt_embeds,
258
+ txt_ids=text_ids,
259
+ txt_mask=prompt_attn_mask, # todo add this
260
+ timesteps=timestep / 1000,
261
+ guidance=guidance
262
+ )
263
+
264
+ if guidance_scale > 1.00001:
265
+ noise_pred_uncond = self.transformer(
266
+ img=latents,
267
+ img_ids=latent_image_ids,
268
+ txt=negative_prompt_embeds,
269
+ txt_ids=negative_text_ids,
270
+ txt_mask=negative_prompt_attn_mask, # todo add this
271
+ timesteps=timestep / 1000,
272
+ guidance=guidance
273
+ )
274
+
275
+ noise_pred = noise_pred_uncond + self.guidance_scale * \
276
+ (noise_pred_text - noise_pred_uncond)
277
+
278
+ else:
279
+ noise_pred = noise_pred_text
280
+
281
+ # compute the previous noisy sample x_t -> x_t-1
282
+ latents_dtype = latents.dtype
283
+ latents = self.scheduler.step(
284
+ noise_pred, t, latents, return_dict=False)[0]
285
+
286
+ if latents.dtype != latents_dtype:
287
+ if torch.backends.mps.is_available():
288
+ # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
289
+ latents = latents.to(latents_dtype)
290
+
291
+ if callback_on_step_end is not None:
292
+ callback_kwargs = {}
293
+ for k in callback_on_step_end_tensor_inputs:
294
+ callback_kwargs[k] = locals()[k]
295
+ callback_outputs = callback_on_step_end(
296
+ self, i, t, callback_kwargs)
297
+
298
+ latents = callback_outputs.pop("latents", latents)
299
+ prompt_embeds = callback_outputs.pop(
300
+ "prompt_embeds", prompt_embeds)
301
+
302
+ # call the callback, if provided
303
+ if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
304
+ progress_bar.update()
305
+
306
+ if XLA_AVAILABLE:
307
+ xm.mark_step()
308
+
309
+ if output_type == "latent":
310
+ image = latents
311
+
312
+ else:
313
+ if not self.is_radiance:
314
+ latents = self._unpack_latents(
315
+ latents, height, width, self.vae_scale_factor)
316
+ latents = (latents / self.vae.config.scaling_factor) + \
317
+ self.vae.config.shift_factor
318
+ image = self.vae.decode(latents, return_dict=False)[0]
319
+ image = self.image_processor.postprocess(
320
+ image, output_type=output_type)
321
+
322
+ # Offload all models
323
+ self.maybe_free_model_hooks()
324
+
325
+ if not return_dict:
326
+ return (image,)
327
+
328
+ return FluxPipelineOutput(images=image)