comdoleger commited on
Commit
eea6325
·
verified ·
1 Parent(s): eba2e24

Upload extensions_built_in/diffusion_models/f_light/src/pipeline.py with huggingface_hub

Browse files
extensions_built_in/diffusion_models/f_light/src/pipeline.py ADDED
@@ -0,0 +1,308 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # originally from https://github.com/fal-ai/f-lite/blob/main/f_lite/pipeline.py but modified slightly
2
+ import logging
3
+ import math
4
+ from dataclasses import dataclass
5
+ from typing import Any, Dict, List, Optional, Tuple, Union
6
+
7
+ import numpy as np
8
+ import torch
9
+ from diffusers import AutoencoderKL, DiffusionPipeline
10
+ from diffusers.utils import BaseOutput
11
+ from diffusers.utils.torch_utils import randn_tensor
12
+ from PIL import Image
13
+ from torch import FloatTensor
14
+ from tqdm.auto import tqdm
15
+ from transformers import T5EncoderModel, T5TokenizerFast
16
+
17
+
18
+
19
+ logger = logging.getLogger(__name__)
20
+
21
+
22
+ @dataclass
23
+ class APGConfig:
24
+ """APG (Augmented Parallel Guidance) configuration"""
25
+
26
+ enabled: bool = True
27
+ orthogonal_threshold: float = 0.03
28
+
29
+
30
+ @dataclass
31
+ class FLitePipelineOutput(BaseOutput):
32
+ """
33
+ Output class for FLitePipeline pipeline.
34
+ Args:
35
+ images (`List[PIL.Image.Image]` or `np.ndarray`)
36
+ List of denoised PIL images of length `batch_size` or numpy array of shape `(batch_size, height, width,
37
+ num_channels)`. PIL images or numpy array present the denoised images of the diffusion pipeline.
38
+ """
39
+
40
+ images: Union[List[Image.Image], np.ndarray]
41
+
42
+
43
+ class FLitePipeline(DiffusionPipeline):
44
+ r"""
45
+ Pipeline for text-to-image generation using F-Lite model.
46
+ This model inherits from [`DiffusionPipeline`].
47
+ """
48
+
49
+ model_cpu_offload_seq = "text_encoder->dit_model->vae"
50
+
51
+ dit_model: torch.nn.Module
52
+ vae: AutoencoderKL
53
+ text_encoder: T5EncoderModel
54
+ tokenizer: T5TokenizerFast
55
+ _progress_bar_config: Dict[str, Any]
56
+
57
+ def __init__(
58
+ self, dit_model: torch.nn.Module, vae: AutoencoderKL, text_encoder: T5EncoderModel, tokenizer: T5TokenizerFast
59
+ ):
60
+ super().__init__()
61
+ # Register all modules for the pipeline
62
+ # Access DiffusionPipeline's register_modules directly to avoid mypy error
63
+ DiffusionPipeline.register_modules(
64
+ self, dit_model=dit_model, vae=vae, text_encoder=text_encoder, tokenizer=tokenizer
65
+ )
66
+
67
+ # Move models to channels last for better performance
68
+ # AutoencoderKL inherits from torch.nn.Module which has these methods
69
+ if hasattr(self.vae, "to"):
70
+ self.vae.to(memory_format=torch.channels_last)
71
+ if hasattr(self.vae, "requires_grad_"):
72
+ self.vae.requires_grad_(False)
73
+ if hasattr(self.text_encoder, "requires_grad_"):
74
+ self.text_encoder.requires_grad_(False)
75
+
76
+ # Constants
77
+ self.vae_scale_factor = 8
78
+ self.return_index = -8 # T5 hidden state index to use
79
+
80
+ def enable_vae_slicing(self):
81
+ """Enable VAE slicing for memory efficiency."""
82
+ if hasattr(self.vae, "enable_slicing"):
83
+ self.vae.enable_slicing()
84
+
85
+ def enable_vae_tiling(self):
86
+ """Enable VAE tiling for memory efficiency."""
87
+ if hasattr(self.vae, "enable_tiling"):
88
+ self.vae.enable_tiling()
89
+
90
+ def set_progress_bar_config(self, **kwargs):
91
+ """Set progress bar configuration."""
92
+ self._progress_bar_config = kwargs
93
+
94
+ def progress_bar(self, iterable=None, **kwargs):
95
+ """Create progress bar for iterations."""
96
+ self._progress_bar_config = getattr(self, "_progress_bar_config", None) or {}
97
+ config = {**self._progress_bar_config, **kwargs}
98
+ return tqdm(iterable, **config)
99
+
100
+ def encode_prompt(
101
+ self,
102
+ prompt: Union[str, List[str]],
103
+ negative_prompt: Optional[Union[str, List[str]]] = None,
104
+ device: Optional[torch.device] = None,
105
+ dtype: Optional[torch.dtype] = None,
106
+ max_sequence_length: int = 512,
107
+ return_index: int = -8,
108
+ ) -> Tuple[FloatTensor, FloatTensor]:
109
+ """Encodes the prompt and negative prompt."""
110
+ if isinstance(prompt, str):
111
+ prompt = [prompt]
112
+ device = device or self.text_encoder.device
113
+ # Text encoder forward pass
114
+ text_inputs = self.tokenizer(
115
+ prompt,
116
+ padding="max_length",
117
+ max_length=max_sequence_length,
118
+ truncation=True,
119
+ return_tensors="pt",
120
+ )
121
+ text_input_ids = text_inputs.input_ids.to(device)
122
+ prompt_embeds = self.text_encoder(text_input_ids, return_dict=True, output_hidden_states=True)
123
+ prompt_embeds_tensor = prompt_embeds.hidden_states[return_index]
124
+ if return_index != -1:
125
+ prompt_embeds_tensor = self.text_encoder.encoder.final_layer_norm(prompt_embeds_tensor)
126
+ prompt_embeds_tensor = self.text_encoder.encoder.dropout(prompt_embeds_tensor)
127
+
128
+ dtype = dtype or next(self.text_encoder.parameters()).dtype
129
+ prompt_embeds_tensor = prompt_embeds_tensor.to(dtype=dtype, device=device)
130
+
131
+ # Handle negative prompts
132
+ if negative_prompt is None:
133
+ negative_embeds = torch.zeros_like(prompt_embeds_tensor)
134
+ else:
135
+ if isinstance(negative_prompt, str):
136
+ negative_prompt = [negative_prompt]
137
+ negative_result = self.encode_prompt(
138
+ prompt=negative_prompt, device=device, dtype=dtype, return_index=return_index
139
+ )
140
+ negative_embeds = negative_result[0]
141
+
142
+ # Explicitly cast both tensors to FloatTensor for mypy
143
+ from typing import cast
144
+
145
+ prompt_tensor = cast(FloatTensor, prompt_embeds_tensor.to(dtype=dtype))
146
+ negative_tensor = cast(FloatTensor, negative_embeds.to(dtype=dtype))
147
+ return (prompt_tensor, negative_tensor)
148
+
149
+ def to(self, torch_device=None, torch_dtype=None, silence_dtype_warnings=False):
150
+ """Move pipeline components to specified device and dtype."""
151
+ if hasattr(self, "vae"):
152
+ self.vae.to(device=torch_device, dtype=torch_dtype)
153
+ if hasattr(self, "text_encoder"):
154
+ self.text_encoder.to(device=torch_device, dtype=torch_dtype)
155
+ if hasattr(self, "dit_model"):
156
+ self.dit_model.to(device=torch_device, dtype=torch_dtype)
157
+ return self
158
+
159
+ @torch.no_grad()
160
+ def __call__(
161
+ self,
162
+ prompt: Union[str, List[str]]=None,
163
+ prompt_embeds: Optional[FloatTensor] = None,
164
+ height: Optional[int] = 1024,
165
+ width: Optional[int] = 1024,
166
+ num_inference_steps: int = 30,
167
+ guidance_scale: float = 6.0,
168
+ negative_prompt: Optional[Union[str, List[str]]] = None,
169
+ negative_prompt_embeds: Optional[FloatTensor] = None,
170
+ num_images_per_prompt: int = 1,
171
+ generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
172
+ dtype: Optional[torch.dtype] = None,
173
+ alpha: Optional[float] = None,
174
+ apg_config: Optional[APGConfig] = None,
175
+ **kwargs,
176
+ ):
177
+ """Generate images from text prompt."""
178
+ # Ensure height and width are not None for calculation
179
+ if height is None:
180
+ height = 1024
181
+ if width is None:
182
+ width = 1024
183
+
184
+ dtype = dtype or next(self.dit_model.parameters()).dtype
185
+ apg_config = apg_config or APGConfig(enabled=False)
186
+
187
+ device = self._execution_device
188
+
189
+ # 2. Encode prompts
190
+ prompt_batch_size = len(prompt) if isinstance(prompt, list) else 1
191
+ batch_size = prompt_batch_size * num_images_per_prompt
192
+
193
+ if prompt_embeds is None or negative_prompt_embeds is None:
194
+ prompt_embeds, negative_embeds = self.encode_prompt(
195
+ prompt=prompt, negative_prompt=negative_prompt, device=self.text_encoder.device, dtype=dtype,
196
+ return_index=self.return_index,
197
+ )
198
+ else:
199
+ negative_embeds = negative_prompt_embeds
200
+
201
+ # Repeat embeddings for num_images_per_prompt
202
+ prompt_embeds = prompt_embeds.repeat_interleave(num_images_per_prompt, dim=0)
203
+ negative_embeds = negative_embeds.repeat_interleave(num_images_per_prompt, dim=0)
204
+
205
+ # 3. Initialize latents
206
+ latent_height = height // self.vae_scale_factor
207
+ latent_width = width // self.vae_scale_factor
208
+
209
+ if isinstance(generator, list):
210
+ if len(generator) != batch_size:
211
+ raise ValueError(f"Got {len(generator)} generators for {batch_size} samples")
212
+
213
+ latents = randn_tensor((batch_size, 16, latent_height, latent_width), generator=generator, device=device, dtype=dtype)
214
+ acc_latents = latents.clone()
215
+
216
+ # 4. Calculate alpha if not provided
217
+ if alpha is None:
218
+ image_token_size = latent_height * latent_width
219
+ alpha = 2 * math.sqrt(image_token_size / (64 * 64))
220
+
221
+ # 6. Sampling loop
222
+ self.dit_model.eval()
223
+
224
+ # Check if guidance is needed
225
+ do_classifier_free_guidance = guidance_scale >= 1.0
226
+
227
+ for i in self.progress_bar(range(num_inference_steps, 0, -1)):
228
+ # Calculate timesteps
229
+ t = i / num_inference_steps
230
+ t_next = (i - 1) / num_inference_steps
231
+ # Scale timesteps according to alpha
232
+ t = t * alpha / (1 + (alpha - 1) * t)
233
+ t_next = t_next * alpha / (1 + (alpha - 1) * t_next)
234
+ dt = t - t_next
235
+
236
+ # Create tensor with proper device
237
+ t_tensor = torch.tensor([t] * batch_size, device=device, dtype=dtype)
238
+
239
+ if do_classifier_free_guidance:
240
+ # Duplicate latents for both conditional and unconditional inputs
241
+ latents_input = torch.cat([latents] * 2)
242
+ # Concatenate negative and positive prompt embeddings
243
+ context_input = torch.cat([negative_embeds, prompt_embeds])
244
+ # Duplicate timesteps for the batch
245
+ t_input = torch.cat([t_tensor] * 2)
246
+
247
+ # Get model predictions in a single pass
248
+ model_outputs = self.dit_model(latents_input, context_input, t_input)
249
+
250
+ # Split outputs back into unconditional and conditional predictions
251
+ uncond_output, cond_output = model_outputs.chunk(2)
252
+
253
+ if apg_config.enabled:
254
+ # Augmented Parallel Guidance
255
+ dy = cond_output
256
+ dd = cond_output - uncond_output
257
+ # Find parallel direction
258
+ parallel_direction = (dy * dd).sum() / (dy * dy).sum() * dy
259
+ orthogonal_direction = dd - parallel_direction
260
+ # Scale orthogonal component
261
+ orthogonal_std = orthogonal_direction.std()
262
+ orthogonal_scale = min(1, apg_config.orthogonal_threshold / orthogonal_std)
263
+ orthogonal_direction = orthogonal_direction * orthogonal_scale
264
+ model_output = dy + (guidance_scale - 1) * orthogonal_direction
265
+ else:
266
+ # Standard classifier-free guidance
267
+ model_output = uncond_output + guidance_scale * (cond_output - uncond_output)
268
+ else:
269
+ # If no guidance needed, just run the model normally
270
+ model_output = self.dit_model(latents, prompt_embeds, t_tensor)
271
+
272
+ # Update latents
273
+ acc_latents = acc_latents + dt * model_output.to(device)
274
+ latents = acc_latents.clone()
275
+
276
+ # 7. Decode latents
277
+ # These checks handle the case where mypy doesn't recognize these attributes
278
+ scaling_factor = getattr(self.vae.config, "scaling_factor", 0.18215) if hasattr(self.vae, "config") else 0.18215
279
+ shift_factor = getattr(self.vae.config, "shift_factor", 0) if hasattr(self.vae, "config") else 0
280
+
281
+ latents = latents / scaling_factor + shift_factor
282
+
283
+ vae_dtype = self.vae.dtype if hasattr(self.vae, "dtype") else dtype
284
+ decoded_images = self.vae.decode(latents.to(vae_dtype)).sample if hasattr(self.vae, "decode") else latents
285
+
286
+ # Offload all models
287
+ try:
288
+ self.maybe_free_model_hooks()
289
+ except AttributeError as e:
290
+ if "OptimizedModule" in str(e):
291
+ import warnings
292
+ warnings.warn(
293
+ "Encountered 'OptimizedModule' error when offloading models. "
294
+ "This issue might be fixed in the future by: "
295
+ "https://github.com/huggingface/diffusers/pull/10730"
296
+ )
297
+ else:
298
+ raise
299
+
300
+ # 8. Post-process images
301
+ images = (decoded_images / 2 + 0.5).clamp(0, 1)
302
+ # Convert to PIL Images
303
+ images = (images * 255).round().clamp(0, 255).to(torch.uint8).cpu()
304
+ pil_images = [Image.fromarray(img.permute(1, 2, 0).numpy()) for img in images]
305
+
306
+ return FLitePipelineOutput(
307
+ images=pil_images,
308
+ )