diffusers-workflow 0.4.0a3__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (171) hide show
  1. diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
  2. diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
  3. diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
  4. diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
  5. diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
  6. diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
  7. dw/__init__.py +353 -0
  8. dw/arguments.py +906 -0
  9. dw/cache_blocks.json +16 -0
  10. dw/cache_blocks.py +145 -0
  11. dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
  12. dw/events.py +78 -0
  13. dw/hub_cache.py +289 -0
  14. dw/introspection.py +458 -0
  15. dw/log_setup.py +45 -0
  16. dw/pipeline_processors/chain.py +750 -0
  17. dw/pipeline_processors/config_objects.py +235 -0
  18. dw/pipeline_processors/pipeline.py +1687 -0
  19. dw/pipeline_processors/remote.py +18 -0
  20. dw/previous_results.py +259 -0
  21. dw/prompt_weighting.py +378 -0
  22. dw/repl.py +298 -0
  23. dw/repl_commands.py +808 -0
  24. dw/repl_worker.py +129 -0
  25. dw/result.py +850 -0
  26. dw/run.py +92 -0
  27. dw/schema.py +24 -0
  28. dw/security.py +379 -0
  29. dw/serve.py +70 -0
  30. dw/server/__init__.py +2 -0
  31. dw/server/app.py +588 -0
  32. dw/server/jobs.py +547 -0
  33. dw/server/ui/assets/abap-08VXUWAP.js +1 -0
  34. dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
  35. dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
  36. dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
  37. dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
  38. dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
  39. dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
  40. dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
  41. dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
  42. dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
  43. dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
  44. dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
  45. dw/server/ui/assets/css-DIMkf-bt.js +3 -0
  46. dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
  47. dw/server/ui/assets/cssMode-CEh6hWi2.js +1 -0
  48. dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
  49. dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
  50. dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
  51. dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
  52. dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
  53. dw/server/ui/assets/editor.api-CExg3_mM.js +847 -0
  54. dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
  55. dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
  56. dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
  57. dw/server/ui/assets/freemarker2-DH6orYh2.js +3 -0
  58. dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
  59. dw/server/ui/assets/go-C-y9NEjX.js +1 -0
  60. dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
  61. dw/server/ui/assets/handlebars-CbrMVW4Q.js +1 -0
  62. dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
  63. dw/server/ui/assets/html-YDNPZw2M.js +1 -0
  64. dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
  65. dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
  66. dw/server/ui/assets/index-B7-VcYS-.css +1 -0
  67. dw/server/ui/assets/index-D_EiPU3b.js +13 -0
  68. dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
  69. dw/server/ui/assets/java-BEtHBSE6.js +1 -0
  70. dw/server/ui/assets/javascript-dYuBvioq.js +1 -0
  71. dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
  72. dw/server/ui/assets/jsonMode-CUqLM39V.js +7 -0
  73. dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
  74. dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
  75. dw/server/ui/assets/less-B9JPFI3C.js +2 -0
  76. dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
  77. dw/server/ui/assets/liquid-D6vxBzMv.js +1 -0
  78. dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.js +4 -0
  79. dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
  80. dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
  81. dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
  82. dw/server/ui/assets/mdx-SHQb6vmD.js +1 -0
  83. dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
  84. dw/server/ui/assets/monaco--ixms01u.css +1 -0
  85. dw/server/ui/assets/monaco-CP-s5rcP.js +56 -0
  86. dw/server/ui/assets/msdax-DauUninz.js +1 -0
  87. dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
  88. dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
  89. dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
  90. dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
  91. dw/server/ui/assets/perl-oz_6vUea.js +1 -0
  92. dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
  93. dw/server/ui/assets/php-nr791fC2.js +1 -0
  94. dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
  95. dw/server/ui/assets/postiats-43DmfD33.js +1 -0
  96. dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
  97. dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
  98. dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
  99. dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
  100. dw/server/ui/assets/python-x0_EGHq9.js +1 -0
  101. dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
  102. dw/server/ui/assets/r-BwWrilGY.js +1 -0
  103. dw/server/ui/assets/razor-BZC4LQDP.js +1 -0
  104. dw/server/ui/assets/redis-ClamHrr6.js +1 -0
  105. dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
  106. dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
  107. dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
  108. dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
  109. dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
  110. dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
  111. dw/server/ui/assets/scheme-BeGwcela.js +1 -0
  112. dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
  113. dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
  114. dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
  115. dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
  116. dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
  117. dw/server/ui/assets/sql-NEE52Syq.js +1 -0
  118. dw/server/ui/assets/st-DbInun42.js +1 -0
  119. dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
  120. dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
  121. dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
  122. dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
  123. dw/server/ui/assets/tsMode-BTfA6SbD.js +11 -0
  124. dw/server/ui/assets/twig-De2hgUGE.js +1 -0
  125. dw/server/ui/assets/typescript-CWA4MsNk.js +1 -0
  126. dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
  127. dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
  128. dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
  129. dw/server/ui/assets/workers-CWU0uvj5.js +1 -0
  130. dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
  131. dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
  132. dw/server/ui/index.html +17 -0
  133. dw/settings.py +77 -0
  134. dw/step.py +132 -0
  135. dw/tasks/audio_utils.py +266 -0
  136. dw/tasks/background_remover.py +43 -0
  137. dw/tasks/borders.py +113 -0
  138. dw/tasks/concat_videos.py +80 -0
  139. dw/tasks/depth_estimator.py +54 -0
  140. dw/tasks/diffusion_upscale.py +109 -0
  141. dw/tasks/format_messages.py +24 -0
  142. dw/tasks/gather.py +139 -0
  143. dw/tasks/image_to_text.py +43 -0
  144. dw/tasks/image_utils.py +661 -0
  145. dw/tasks/interpolate_frames.py +227 -0
  146. dw/tasks/model_cache.py +39 -0
  147. dw/tasks/pair_audio.py +58 -0
  148. dw/tasks/qr_code.py +19 -0
  149. dw/tasks/restore_faces.py +175 -0
  150. dw/tasks/rife_model.py +192 -0
  151. dw/tasks/segment.py +121 -0
  152. dw/tasks/task.py +474 -0
  153. dw/tasks/tensor_image.py +57 -0
  154. dw/tasks/text_generation.py +168 -0
  155. dw/tasks/text_sections.py +80 -0
  156. dw/tasks/upscale.py +203 -0
  157. dw/tasks/video_utils.py +154 -0
  158. dw/tasks/zoe_depth.py +71 -0
  159. dw/teacache.py +376 -0
  160. dw/teacache_models.json +99 -0
  161. dw/test.py +29 -0
  162. dw/type_helpers.py +68 -0
  163. dw/validate.py +43 -0
  164. dw/variables.py +153 -0
  165. dw/worker.py +517 -0
  166. dw/workflow.py +553 -0
  167. dw/workflow_schema.json +1157 -0
  168. dw/workflows/augment_prompt.json +65 -0
  169. dw/workflows/describe_image.json +58 -0
  170. dw/workflows/h3_context_ir.json +57 -0
  171. dw/workflows/test.json +31 -0
@@ -0,0 +1,1184 @@
1
+ # Copyright 2024 Black Forest Labs and The HuggingFace Team. All rights reserved.
2
+ # modeled after RF Inversion: https://rf-inversion.github.io/, authored by Litu Rout, Yujia Chen, Nataniel Ruiz,
3
+ # Constantine Caramanis, Sanjay Shakkottai and Wen-Sheng Chu.
4
+ #
5
+ # Licensed under the Apache License, Version 2.0 (the "License");
6
+ # you may not use this file except in compliance with the License.
7
+ # You may obtain a copy of the License at
8
+ #
9
+ # http://www.apache.org/licenses/LICENSE-2.0
10
+ #
11
+ # Unless required by applicable law or agreed to in writing, software
12
+ # distributed under the License is distributed on an "AS IS" BASIS,
13
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14
+ # See the License for the specific language governing permissions and
15
+ # limitations under the License.
16
+
17
+ import inspect
18
+ from typing import Any, Callable, Dict, List, Optional, Union
19
+
20
+ import numpy as np
21
+ import torch
22
+ from transformers import CLIPTextModel, CLIPTokenizer, T5EncoderModel, T5TokenizerFast
23
+
24
+ from diffusers.image_processor import PipelineImageInput, VaeImageProcessor
25
+ from diffusers.loaders import (
26
+ FluxLoraLoaderMixin,
27
+ FromSingleFileMixin,
28
+ TextualInversionLoaderMixin,
29
+ )
30
+ from diffusers.models.autoencoders import AutoencoderKL
31
+ from diffusers.models.transformers import FluxTransformer2DModel
32
+ from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput
33
+ from diffusers.pipelines.pipeline_utils import DiffusionPipeline
34
+ from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
35
+ from diffusers.utils import (
36
+ USE_PEFT_BACKEND,
37
+ is_torch_xla_available,
38
+ logging,
39
+ replace_example_docstring,
40
+ scale_lora_layers,
41
+ unscale_lora_layers,
42
+ )
43
+ from diffusers.utils.torch_utils import randn_tensor
44
+
45
+ if is_torch_xla_available():
46
+ import torch_xla.core.xla_model as xm
47
+
48
+ XLA_AVAILABLE = True
49
+ else:
50
+ XLA_AVAILABLE = False
51
+
52
+
53
+ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
54
+
55
+ EXAMPLE_DOC_STRING = """
56
+ Examples:
57
+ ```py
58
+ >>> import torch
59
+ >>> import requests
60
+ >>> import PIL
61
+ >>> from io import BytesIO
62
+ >>> from diffusers import DiffusionPipeline
63
+
64
+ >>> pipe = DiffusionPipeline.from_pretrained(
65
+ ... "black-forest-labs/FLUX.1-dev",
66
+ ... torch_dtype=torch.bfloat16,
67
+ ... custom_pipeline="pipeline_flux_rf_inversion")
68
+ >>> pipe.to("cuda")
69
+
70
+ >>> def download_image(url):
71
+ ... response = requests.get(url)
72
+ ... return PIL.Image.open(BytesIO(response.content)).convert("RGB")
73
+
74
+
75
+ >>> img_url = "https://www.aiml.informatik.tu-darmstadt.de/people/mbrack/tennis.jpg"
76
+ >>> image = download_image(img_url)
77
+
78
+ >>> inverted_latents, image_latents, latent_image_ids = pipe.invert(image=image, num_inversion_steps=28, gamma=0.5)
79
+
80
+ >>> edited_image = pipe(
81
+ ... prompt="a tomato",
82
+ ... inverted_latents=inverted_latents,
83
+ ... image_latents=image_latents,
84
+ ... latent_image_ids=latent_image_ids,
85
+ ... start_timestep=0,
86
+ ... stop_timestep=.25,
87
+ ... num_inference_steps=28,
88
+ ... eta=0.9,
89
+ ... ).images[0]
90
+ ```
91
+ """
92
+
93
+
94
+ # Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shift
95
+ def calculate_shift(
96
+ image_seq_len,
97
+ base_seq_len: int = 256,
98
+ max_seq_len: int = 4096,
99
+ base_shift: float = 0.5,
100
+ max_shift: float = 1.16,
101
+ ):
102
+ m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
103
+ b = base_shift - m * base_seq_len
104
+ mu = image_seq_len * m + b
105
+ return mu
106
+
107
+
108
+ # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
109
+ def retrieve_timesteps(
110
+ scheduler,
111
+ num_inference_steps: Optional[int] = None,
112
+ device: Optional[Union[str, torch.device]] = None,
113
+ timesteps: Optional[List[int]] = None,
114
+ sigmas: Optional[List[float]] = None,
115
+ **kwargs,
116
+ ):
117
+ r"""
118
+ Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
119
+ custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
120
+
121
+ Args:
122
+ scheduler (`SchedulerMixin`):
123
+ The scheduler to get timesteps from.
124
+ num_inference_steps (`int`):
125
+ The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
126
+ must be `None`.
127
+ device (`str` or `torch.device`, *optional*):
128
+ The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
129
+ timesteps (`List[int]`, *optional*):
130
+ Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
131
+ `num_inference_steps` and `sigmas` must be `None`.
132
+ sigmas (`List[float]`, *optional*):
133
+ Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
134
+ `num_inference_steps` and `timesteps` must be `None`.
135
+
136
+ Returns:
137
+ `Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
138
+ second element is the number of inference steps.
139
+ """
140
+ if timesteps is not None and sigmas is not None:
141
+ raise ValueError(
142
+ "Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values"
143
+ )
144
+ if timesteps is not None:
145
+ accepts_timesteps = "timesteps" in set(
146
+ inspect.signature(scheduler.set_timesteps).parameters.keys()
147
+ )
148
+ if not accepts_timesteps:
149
+ raise ValueError(
150
+ f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
151
+ f" timestep schedules. Please check whether you are using the correct scheduler."
152
+ )
153
+ scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
154
+ timesteps = scheduler.timesteps
155
+ num_inference_steps = len(timesteps)
156
+ elif sigmas is not None:
157
+ accept_sigmas = "sigmas" in set(
158
+ inspect.signature(scheduler.set_timesteps).parameters.keys()
159
+ )
160
+ if not accept_sigmas:
161
+ raise ValueError(
162
+ f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
163
+ f" sigmas schedules. Please check whether you are using the correct scheduler."
164
+ )
165
+ scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
166
+ timesteps = scheduler.timesteps
167
+ num_inference_steps = len(timesteps)
168
+ else:
169
+ scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
170
+ timesteps = scheduler.timesteps
171
+ return timesteps, num_inference_steps
172
+
173
+
174
+ class RFInversionFluxPipeline(
175
+ DiffusionPipeline,
176
+ FluxLoraLoaderMixin,
177
+ FromSingleFileMixin,
178
+ TextualInversionLoaderMixin,
179
+ ):
180
+ r"""
181
+ The Flux pipeline for text-to-image generation.
182
+
183
+ Reference: https://blackforestlabs.ai/announcing-black-forest-labs/
184
+
185
+ Args:
186
+ transformer ([`FluxTransformer2DModel`]):
187
+ Conditional Transformer (MMDiT) architecture to denoise the encoded image latents.
188
+ scheduler ([`FlowMatchEulerDiscreteScheduler`]):
189
+ A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
190
+ vae ([`AutoencoderKL`]):
191
+ Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
192
+ text_encoder ([`CLIPTextModel`]):
193
+ [CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically
194
+ the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant.
195
+ text_encoder_2 ([`T5EncoderModel`]):
196
+ [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
197
+ the [google/t5-v1_1-xxl](https://huggingface.co/google/t5-v1_1-xxl) variant.
198
+ tokenizer (`CLIPTokenizer`):
199
+ Tokenizer of class
200
+ [CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
201
+ tokenizer_2 (`T5TokenizerFast`):
202
+ Second Tokenizer of class
203
+ [T5TokenizerFast](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5TokenizerFast).
204
+ """
205
+
206
+ model_cpu_offload_seq = "text_encoder->text_encoder_2->transformer->vae"
207
+ _optional_components = []
208
+ _callback_tensor_inputs = ["latents", "prompt_embeds"]
209
+
210
+ def __init__(
211
+ self,
212
+ scheduler: FlowMatchEulerDiscreteScheduler,
213
+ vae: AutoencoderKL,
214
+ text_encoder: CLIPTextModel,
215
+ tokenizer: CLIPTokenizer,
216
+ text_encoder_2: T5EncoderModel,
217
+ tokenizer_2: T5TokenizerFast,
218
+ transformer: FluxTransformer2DModel,
219
+ ):
220
+ super().__init__()
221
+
222
+ self.register_modules(
223
+ vae=vae,
224
+ text_encoder=text_encoder,
225
+ text_encoder_2=text_encoder_2,
226
+ tokenizer=tokenizer,
227
+ tokenizer_2=tokenizer_2,
228
+ transformer=transformer,
229
+ scheduler=scheduler,
230
+ )
231
+ self.vae_scale_factor = (
232
+ 2 ** (len(self.vae.config.block_out_channels) - 1)
233
+ if hasattr(self, "vae") and self.vae is not None
234
+ else 8
235
+ )
236
+ self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
237
+ self.tokenizer_max_length = (
238
+ self.tokenizer.model_max_length
239
+ if hasattr(self, "tokenizer") and self.tokenizer is not None
240
+ else 77
241
+ )
242
+ self.default_sample_size = 128
243
+
244
+ # Copied from diffusers.pipelines.flux.pipeline_flux.FluxPipeline._get_t5_prompt_embeds
245
+ def _get_t5_prompt_embeds(
246
+ self,
247
+ prompt: Union[str, List[str]] = None,
248
+ num_images_per_prompt: int = 1,
249
+ max_sequence_length: int = 512,
250
+ device: Optional[torch.device] = None,
251
+ dtype: Optional[torch.dtype] = None,
252
+ ):
253
+ device = device or self._execution_device
254
+ dtype = dtype or self.text_encoder.dtype
255
+
256
+ prompt = [prompt] if isinstance(prompt, str) else prompt
257
+ batch_size = len(prompt)
258
+
259
+ if isinstance(self, TextualInversionLoaderMixin):
260
+ prompt = self.maybe_convert_prompt(prompt, self.tokenizer_2)
261
+
262
+ text_inputs = self.tokenizer_2(
263
+ prompt,
264
+ padding="max_length",
265
+ max_length=max_sequence_length,
266
+ truncation=True,
267
+ return_length=False,
268
+ return_overflowing_tokens=False,
269
+ return_tensors="pt",
270
+ )
271
+ text_input_ids = text_inputs.input_ids
272
+ untruncated_ids = self.tokenizer_2(
273
+ prompt, padding="longest", return_tensors="pt"
274
+ ).input_ids
275
+
276
+ if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
277
+ text_input_ids, untruncated_ids
278
+ ):
279
+ removed_text = self.tokenizer_2.batch_decode(
280
+ untruncated_ids[:, self.tokenizer_max_length - 1 : -1]
281
+ )
282
+ logger.warning(
283
+ "The following part of your input was truncated because `max_sequence_length` is set to "
284
+ f" {max_sequence_length} tokens: {removed_text}"
285
+ )
286
+
287
+ prompt_embeds = self.text_encoder_2(
288
+ text_input_ids.to(device), output_hidden_states=False
289
+ )[0]
290
+
291
+ dtype = self.text_encoder_2.dtype
292
+ prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
293
+
294
+ _, seq_len, _ = prompt_embeds.shape
295
+
296
+ # duplicate text embeddings and attention mask for each generation per prompt, using mps friendly method
297
+ prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
298
+ prompt_embeds = prompt_embeds.view(
299
+ batch_size * num_images_per_prompt, seq_len, -1
300
+ )
301
+
302
+ return prompt_embeds
303
+
304
+ # Copied from diffusers.pipelines.flux.pipeline_flux.FluxPipeline._get_clip_prompt_embeds
305
+ def _get_clip_prompt_embeds(
306
+ self,
307
+ prompt: Union[str, List[str]],
308
+ num_images_per_prompt: int = 1,
309
+ device: Optional[torch.device] = None,
310
+ ):
311
+ device = device or self._execution_device
312
+
313
+ prompt = [prompt] if isinstance(prompt, str) else prompt
314
+ batch_size = len(prompt)
315
+
316
+ if isinstance(self, TextualInversionLoaderMixin):
317
+ prompt = self.maybe_convert_prompt(prompt, self.tokenizer)
318
+
319
+ text_inputs = self.tokenizer(
320
+ prompt,
321
+ padding="max_length",
322
+ max_length=self.tokenizer_max_length,
323
+ truncation=True,
324
+ return_overflowing_tokens=False,
325
+ return_length=False,
326
+ return_tensors="pt",
327
+ )
328
+
329
+ text_input_ids = text_inputs.input_ids
330
+ untruncated_ids = self.tokenizer(
331
+ prompt, padding="longest", return_tensors="pt"
332
+ ).input_ids
333
+ if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
334
+ text_input_ids, untruncated_ids
335
+ ):
336
+ removed_text = self.tokenizer.batch_decode(
337
+ untruncated_ids[:, self.tokenizer_max_length - 1 : -1]
338
+ )
339
+ logger.warning(
340
+ "The following part of your input was truncated because CLIP can only handle sequences up to"
341
+ f" {self.tokenizer_max_length} tokens: {removed_text}"
342
+ )
343
+ prompt_embeds = self.text_encoder(
344
+ text_input_ids.to(device), output_hidden_states=False
345
+ )
346
+
347
+ # Use pooled output of CLIPTextModel
348
+ prompt_embeds = prompt_embeds.pooler_output
349
+ prompt_embeds = prompt_embeds.to(dtype=self.text_encoder.dtype, device=device)
350
+
351
+ # duplicate text embeddings for each generation per prompt, using mps friendly method
352
+ prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt)
353
+ prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, -1)
354
+
355
+ return prompt_embeds
356
+
357
+ # Copied from diffusers.pipelines.flux.pipeline_flux.FluxPipeline.encode_prompt
358
+ def encode_prompt(
359
+ self,
360
+ prompt: Union[str, List[str]],
361
+ prompt_2: Union[str, List[str]],
362
+ device: Optional[torch.device] = None,
363
+ num_images_per_prompt: int = 1,
364
+ prompt_embeds: Optional[torch.FloatTensor] = None,
365
+ pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
366
+ max_sequence_length: int = 512,
367
+ lora_scale: Optional[float] = None,
368
+ ):
369
+ r"""
370
+
371
+ Args:
372
+ prompt (`str` or `List[str]`, *optional*):
373
+ prompt to be encoded
374
+ prompt_2 (`str` or `List[str]`, *optional*):
375
+ The prompt or prompts to be sent to the `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
376
+ used in all text-encoders
377
+ device: (`torch.device`):
378
+ torch device
379
+ num_images_per_prompt (`int`):
380
+ number of images that should be generated per prompt
381
+ prompt_embeds (`torch.FloatTensor`, *optional*):
382
+ Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
383
+ provided, text embeddings will be generated from `prompt` input argument.
384
+ pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
385
+ Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
386
+ If not provided, pooled text embeddings will be generated from `prompt` input argument.
387
+ lora_scale (`float`, *optional*):
388
+ A lora scale that will be applied to all LoRA layers of the text encoder if LoRA layers are loaded.
389
+ """
390
+ device = device or self._execution_device
391
+
392
+ # set lora scale so that monkey patched LoRA
393
+ # function of text encoder can correctly access it
394
+ if lora_scale is not None and isinstance(self, FluxLoraLoaderMixin):
395
+ self._lora_scale = lora_scale
396
+
397
+ # dynamically adjust the LoRA scale
398
+ if self.text_encoder is not None and USE_PEFT_BACKEND:
399
+ scale_lora_layers(self.text_encoder, lora_scale)
400
+ if self.text_encoder_2 is not None and USE_PEFT_BACKEND:
401
+ scale_lora_layers(self.text_encoder_2, lora_scale)
402
+
403
+ prompt = [prompt] if isinstance(prompt, str) else prompt
404
+
405
+ if prompt_embeds is None:
406
+ prompt_2 = prompt_2 or prompt
407
+ prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2
408
+
409
+ # We only use the pooled prompt output from the CLIPTextModel
410
+ pooled_prompt_embeds = self._get_clip_prompt_embeds(
411
+ prompt=prompt,
412
+ device=device,
413
+ num_images_per_prompt=num_images_per_prompt,
414
+ )
415
+ prompt_embeds = self._get_t5_prompt_embeds(
416
+ prompt=prompt_2,
417
+ num_images_per_prompt=num_images_per_prompt,
418
+ max_sequence_length=max_sequence_length,
419
+ device=device,
420
+ )
421
+
422
+ if self.text_encoder is not None:
423
+ if isinstance(self, FluxLoraLoaderMixin) and USE_PEFT_BACKEND:
424
+ # Retrieve the original scale by scaling back the LoRA layers
425
+ unscale_lora_layers(self.text_encoder, lora_scale)
426
+
427
+ if self.text_encoder_2 is not None:
428
+ if isinstance(self, FluxLoraLoaderMixin) and USE_PEFT_BACKEND:
429
+ # Retrieve the original scale by scaling back the LoRA layers
430
+ unscale_lora_layers(self.text_encoder_2, lora_scale)
431
+
432
+ dtype = (
433
+ self.text_encoder.dtype
434
+ if self.text_encoder is not None
435
+ else self.transformer.dtype
436
+ )
437
+ text_ids = torch.zeros(prompt_embeds.shape[1], 3).to(device=device, dtype=dtype)
438
+
439
+ return prompt_embeds, pooled_prompt_embeds, text_ids
440
+
441
+ @torch.no_grad()
442
+ # Modified from diffusers.pipelines.ledits_pp.pipeline_leditspp_stable_diffusion.LEditsPPPipelineStableDiffusion.encode_image
443
+ def encode_image(
444
+ self,
445
+ image,
446
+ dtype=None,
447
+ height=None,
448
+ width=None,
449
+ resize_mode="default",
450
+ crops_coords=None,
451
+ ):
452
+ image = self.image_processor.preprocess(
453
+ image=image,
454
+ height=height,
455
+ width=width,
456
+ resize_mode=resize_mode,
457
+ crops_coords=crops_coords,
458
+ )
459
+ resized = self.image_processor.postprocess(image=image, output_type="pil")
460
+
461
+ if max(image.shape[-2:]) > self.vae.config["sample_size"] * 1.5:
462
+ logger.warning(
463
+ "Your input images far exceed the default resolution of the underlying diffusion model. "
464
+ "The output images may contain severe artifacts! "
465
+ "Consider down-sampling the input using the `height` and `width` parameters"
466
+ )
467
+ image = image.to(dtype)
468
+
469
+ x0 = self.vae.encode(image.to(self.device)).latent_dist.sample()
470
+ x0 = (x0 - self.vae.config.shift_factor) * self.vae.config.scaling_factor
471
+ x0 = x0.to(dtype)
472
+ return x0, resized
473
+
474
+ def check_inputs(
475
+ self,
476
+ prompt,
477
+ prompt_2,
478
+ inverted_latents,
479
+ image_latents,
480
+ latent_image_ids,
481
+ height,
482
+ width,
483
+ start_timestep,
484
+ stop_timestep,
485
+ prompt_embeds=None,
486
+ pooled_prompt_embeds=None,
487
+ callback_on_step_end_tensor_inputs=None,
488
+ max_sequence_length=None,
489
+ ):
490
+ if height % self.vae_scale_factor != 0 or width % self.vae_scale_factor != 0:
491
+ raise ValueError(
492
+ f"`height` and `width` have to be divisible by {self.vae_scale_factor} but are {height} and {width}."
493
+ )
494
+
495
+ if callback_on_step_end_tensor_inputs is not None and not all(
496
+ k in self._callback_tensor_inputs
497
+ for k in callback_on_step_end_tensor_inputs
498
+ ):
499
+ raise ValueError(
500
+ f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
501
+ )
502
+
503
+ if prompt is not None and prompt_embeds is not None:
504
+ raise ValueError(
505
+ f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
506
+ " only forward one of the two."
507
+ )
508
+ elif prompt_2 is not None and prompt_embeds is not None:
509
+ raise ValueError(
510
+ f"Cannot forward both `prompt_2`: {prompt_2} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
511
+ " only forward one of the two."
512
+ )
513
+ elif prompt is None and prompt_embeds is None:
514
+ raise ValueError(
515
+ "Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
516
+ )
517
+ elif prompt is not None and (
518
+ not isinstance(prompt, str) and not isinstance(prompt, list)
519
+ ):
520
+ raise ValueError(
521
+ f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"
522
+ )
523
+ elif prompt_2 is not None and (
524
+ not isinstance(prompt_2, str) and not isinstance(prompt_2, list)
525
+ ):
526
+ raise ValueError(
527
+ f"`prompt_2` has to be of type `str` or `list` but is {type(prompt_2)}"
528
+ )
529
+
530
+ if prompt_embeds is not None and pooled_prompt_embeds is None:
531
+ raise ValueError(
532
+ "If `prompt_embeds` are provided, `pooled_prompt_embeds` also have to be passed. Make sure to generate `pooled_prompt_embeds` from the same text encoder that was used to generate `prompt_embeds`."
533
+ )
534
+
535
+ if max_sequence_length is not None and max_sequence_length > 512:
536
+ raise ValueError(
537
+ f"`max_sequence_length` cannot be greater than 512 but is {max_sequence_length}"
538
+ )
539
+
540
+ if inverted_latents is not None and (
541
+ image_latents is None or latent_image_ids is None
542
+ ):
543
+ raise ValueError(
544
+ "If `inverted_latents` are provided, `image_latents` and `latent_image_ids` also have to be passed. "
545
+ )
546
+ # check start_timestep and stop_timestep
547
+ if start_timestep < 0 or start_timestep > stop_timestep:
548
+ raise ValueError(
549
+ f"`start_timestep` should be in [0, stop_timestep] but is {start_timestep}"
550
+ )
551
+
552
+ @staticmethod
553
+ def _prepare_latent_image_ids(batch_size, height, width, device, dtype):
554
+ latent_image_ids = torch.zeros(height, width, 3)
555
+ latent_image_ids[..., 1] = (
556
+ latent_image_ids[..., 1] + torch.arange(height)[:, None]
557
+ )
558
+ latent_image_ids[..., 2] = (
559
+ latent_image_ids[..., 2] + torch.arange(width)[None, :]
560
+ )
561
+
562
+ latent_image_id_height, latent_image_id_width, latent_image_id_channels = (
563
+ latent_image_ids.shape
564
+ )
565
+
566
+ latent_image_ids = latent_image_ids.reshape(
567
+ latent_image_id_height * latent_image_id_width, latent_image_id_channels
568
+ )
569
+
570
+ return latent_image_ids.to(device=device, dtype=dtype)
571
+
572
+ @staticmethod
573
+ def _pack_latents(latents, batch_size, num_channels_latents, height, width):
574
+ latents = latents.view(
575
+ batch_size, num_channels_latents, height // 2, 2, width // 2, 2
576
+ )
577
+ latents = latents.permute(0, 2, 4, 1, 3, 5)
578
+ latents = latents.reshape(
579
+ batch_size, (height // 2) * (width // 2), num_channels_latents * 4
580
+ )
581
+
582
+ return latents
583
+
584
+ @staticmethod
585
+ def _unpack_latents(latents, height, width, vae_scale_factor):
586
+ batch_size, num_patches, channels = latents.shape
587
+
588
+ height = height // vae_scale_factor
589
+ width = width // vae_scale_factor
590
+
591
+ latents = latents.view(batch_size, height // 2, width // 2, channels // 4, 2, 2)
592
+ latents = latents.permute(0, 3, 1, 4, 2, 5)
593
+
594
+ latents = latents.reshape(batch_size, channels // (2 * 2), height, width)
595
+
596
+ return latents
597
+
598
+ def enable_vae_slicing(self):
599
+ r"""
600
+ Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
601
+ compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
602
+ """
603
+ self.vae.enable_slicing()
604
+
605
+ def disable_vae_slicing(self):
606
+ r"""
607
+ Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to
608
+ computing decoding in one step.
609
+ """
610
+ self.vae.disable_slicing()
611
+
612
+ def enable_vae_tiling(self):
613
+ r"""
614
+ Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
615
+ compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
616
+ processing larger images.
617
+ """
618
+ self.vae.enable_tiling()
619
+
620
+ def disable_vae_tiling(self):
621
+ r"""
622
+ Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to
623
+ computing decoding in one step.
624
+ """
625
+ self.vae.disable_tiling()
626
+
627
+ def prepare_latents_inversion(
628
+ self,
629
+ batch_size,
630
+ num_channels_latents,
631
+ height,
632
+ width,
633
+ dtype,
634
+ device,
635
+ image_latents,
636
+ ):
637
+ height = int(height) // self.vae_scale_factor
638
+ width = int(width) // self.vae_scale_factor
639
+
640
+ latents = self._pack_latents(
641
+ image_latents, batch_size, num_channels_latents, height, width
642
+ )
643
+
644
+ latent_image_ids = self._prepare_latent_image_ids(
645
+ batch_size, height // 2, width // 2, device, dtype
646
+ )
647
+
648
+ return latents, latent_image_ids
649
+
650
+ # Copied from diffusers.pipelines.flux.pipeline_flux.FluxPipeline.prepare_latents
651
+ def prepare_latents(
652
+ self,
653
+ batch_size,
654
+ num_channels_latents,
655
+ height,
656
+ width,
657
+ dtype,
658
+ device,
659
+ generator,
660
+ latents=None,
661
+ ):
662
+ # VAE applies 8x compression on images but we must also account for packing which requires
663
+ # latent height and width to be divisible by 2.
664
+ height = 2 * (int(height) // (self.vae_scale_factor * 2))
665
+ width = 2 * (int(width) // (self.vae_scale_factor * 2))
666
+
667
+ shape = (batch_size, num_channels_latents, height, width)
668
+
669
+ if latents is not None:
670
+ latent_image_ids = self._prepare_latent_image_ids(
671
+ batch_size, height // 2, width // 2, device, dtype
672
+ )
673
+ return latents.to(device=device, dtype=dtype), latent_image_ids
674
+
675
+ if isinstance(generator, list) and len(generator) != batch_size:
676
+ raise ValueError(
677
+ f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
678
+ f" size of {batch_size}. Make sure the batch size matches the length of the generators."
679
+ )
680
+
681
+ latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
682
+ latents = self._pack_latents(
683
+ latents, batch_size, num_channels_latents, height, width
684
+ )
685
+
686
+ latent_image_ids = self._prepare_latent_image_ids(
687
+ batch_size, height // 2, width // 2, device, dtype
688
+ )
689
+
690
+ return latents, latent_image_ids
691
+
692
+ # Copied from diffusers.pipelines.stable_diffusion_3.pipeline_stable_diffusion_3_img2img.StableDiffusion3Img2ImgPipeline.get_timesteps
693
+ def get_timesteps(self, num_inference_steps, strength=1.0):
694
+ # get the original timestep using init_timestep
695
+ init_timestep = min(num_inference_steps * strength, num_inference_steps)
696
+
697
+ t_start = int(max(num_inference_steps - init_timestep, 0))
698
+ timesteps = self.scheduler.timesteps[t_start * self.scheduler.order :]
699
+ sigmas = self.scheduler.sigmas[t_start * self.scheduler.order :]
700
+ if hasattr(self.scheduler, "set_begin_index"):
701
+ self.scheduler.set_begin_index(t_start * self.scheduler.order)
702
+
703
+ return timesteps, sigmas, num_inference_steps - t_start
704
+
705
+ @property
706
+ def guidance_scale(self):
707
+ return self._guidance_scale
708
+
709
+ @property
710
+ def joint_attention_kwargs(self):
711
+ return self._joint_attention_kwargs
712
+
713
+ @property
714
+ def num_timesteps(self):
715
+ return self._num_timesteps
716
+
717
+ @property
718
+ def interrupt(self):
719
+ return self._interrupt
720
+
721
+ @torch.no_grad()
722
+ @replace_example_docstring(EXAMPLE_DOC_STRING)
723
+ def __call__(
724
+ self,
725
+ prompt: Union[str, List[str]] = None,
726
+ prompt_2: Optional[Union[str, List[str]]] = None,
727
+ inverted_latents: Optional[torch.FloatTensor] = None,
728
+ image_latents: Optional[torch.FloatTensor] = None,
729
+ latent_image_ids: Optional[torch.FloatTensor] = None,
730
+ height: Optional[int] = None,
731
+ width: Optional[int] = None,
732
+ eta: float = 1.0,
733
+ decay_eta: Optional[bool] = False,
734
+ eta_decay_power: Optional[float] = 1.0,
735
+ strength: float = 1.0,
736
+ start_timestep: float = 0,
737
+ stop_timestep: float = 0.25,
738
+ num_inference_steps: int = 28,
739
+ sigmas: Optional[List[float]] = None,
740
+ timesteps: List[int] = None,
741
+ guidance_scale: float = 3.5,
742
+ num_images_per_prompt: Optional[int] = 1,
743
+ generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
744
+ latents: Optional[torch.FloatTensor] = None,
745
+ prompt_embeds: Optional[torch.FloatTensor] = None,
746
+ pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
747
+ output_type: Optional[str] = "pil",
748
+ return_dict: bool = True,
749
+ joint_attention_kwargs: Optional[Dict[str, Any]] = None,
750
+ callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
751
+ callback_on_step_end_tensor_inputs: List[str] = ["latents"],
752
+ max_sequence_length: int = 512,
753
+ ):
754
+ r"""
755
+ Function invoked when calling the pipeline for generation.
756
+
757
+ Args:
758
+ prompt (`str` or `List[str]`, *optional*):
759
+ The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
760
+ instead.
761
+ prompt_2 (`str` or `List[str]`, *optional*):
762
+ The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
763
+ will be used instead
764
+ inverted_latents (`torch.Tensor`, *optional*):
765
+ The inverted latents from `pipe.invert`.
766
+ image_latents (`torch.Tensor`, *optional*):
767
+ The image latents from `pipe.invert`.
768
+ latent_image_ids (`torch.Tensor`, *optional*):
769
+ The latent image ids from `pipe.invert`.
770
+ height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
771
+ The height in pixels of the generated image. This is set to 1024 by default for the best results.
772
+ width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
773
+ The width in pixels of the generated image. This is set to 1024 by default for the best results.
774
+ eta (`float`, *optional*, defaults to 1.0):
775
+ The controller guidance, balancing faithfulness & editability:
776
+ higher eta - better faithfullness, less editability. For more significant edits, lower the value of eta.
777
+ num_inference_steps (`int`, *optional*, defaults to 50):
778
+ The number of denoising steps. More denoising steps usually lead to a higher quality image at the
779
+ expense of slower inference.
780
+ timesteps (`List[int]`, *optional*):
781
+ Custom timesteps to use for the denoising process with schedulers which support a `timesteps` argument
782
+ in their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is
783
+ passed will be used. Must be in descending order.
784
+ guidance_scale (`float`, *optional*, defaults to 7.0):
785
+ Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
786
+ `guidance_scale` is defined as `w` of equation 2. of [Imagen
787
+ Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
788
+ 1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
789
+ usually at the expense of lower image quality.
790
+ num_images_per_prompt (`int`, *optional*, defaults to 1):
791
+ The number of images to generate per prompt.
792
+ generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
793
+ One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
794
+ to make generation deterministic.
795
+ latents (`torch.FloatTensor`, *optional*):
796
+ Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
797
+ generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
798
+ tensor will ge generated by sampling using the supplied random `generator`.
799
+ prompt_embeds (`torch.FloatTensor`, *optional*):
800
+ Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
801
+ provided, text embeddings will be generated from `prompt` input argument.
802
+ pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
803
+ Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
804
+ If not provided, pooled text embeddings will be generated from `prompt` input argument.
805
+ output_type (`str`, *optional*, defaults to `"pil"`):
806
+ The output format of the generate image. Choose between
807
+ [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
808
+ return_dict (`bool`, *optional*, defaults to `True`):
809
+ Whether to return a [`~pipelines.flux.FluxPipelineOutput`] instead of a plain tuple.
810
+ joint_attention_kwargs (`dict`, *optional*):
811
+ A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
812
+ `self.processor` in
813
+ [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
814
+ callback_on_step_end (`Callable`, *optional*):
815
+ A function that calls at the end of each denoising steps during the inference. The function is called
816
+ with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
817
+ callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
818
+ `callback_on_step_end_tensor_inputs`.
819
+ callback_on_step_end_tensor_inputs (`List`, *optional*):
820
+ The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
821
+ will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
822
+ `._callback_tensor_inputs` attribute of your pipeline class.
823
+ max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`.
824
+
825
+ Examples:
826
+
827
+ Returns:
828
+ [`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict`
829
+ is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated
830
+ images.
831
+ """
832
+
833
+ height = height or self.default_sample_size * self.vae_scale_factor
834
+ width = width or self.default_sample_size * self.vae_scale_factor
835
+
836
+ # 1. Check inputs. Raise error if not correct
837
+ self.check_inputs(
838
+ prompt,
839
+ prompt_2,
840
+ inverted_latents,
841
+ image_latents,
842
+ latent_image_ids,
843
+ height,
844
+ width,
845
+ start_timestep,
846
+ stop_timestep,
847
+ prompt_embeds=prompt_embeds,
848
+ pooled_prompt_embeds=pooled_prompt_embeds,
849
+ callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
850
+ max_sequence_length=max_sequence_length,
851
+ )
852
+
853
+ self._guidance_scale = guidance_scale
854
+ self._joint_attention_kwargs = joint_attention_kwargs
855
+ self._interrupt = False
856
+ do_rf_inversion = inverted_latents is not None
857
+
858
+ # 2. Define call parameters
859
+ if prompt is not None and isinstance(prompt, str):
860
+ batch_size = 1
861
+ elif prompt is not None and isinstance(prompt, list):
862
+ batch_size = len(prompt)
863
+ else:
864
+ batch_size = prompt_embeds.shape[0]
865
+
866
+ device = self._execution_device
867
+
868
+ lora_scale = (
869
+ self.joint_attention_kwargs.get("scale", None)
870
+ if self.joint_attention_kwargs is not None
871
+ else None
872
+ )
873
+ (
874
+ prompt_embeds,
875
+ pooled_prompt_embeds,
876
+ text_ids,
877
+ ) = self.encode_prompt(
878
+ prompt=prompt,
879
+ prompt_2=prompt_2,
880
+ prompt_embeds=prompt_embeds,
881
+ pooled_prompt_embeds=pooled_prompt_embeds,
882
+ device=device,
883
+ num_images_per_prompt=num_images_per_prompt,
884
+ max_sequence_length=max_sequence_length,
885
+ lora_scale=lora_scale,
886
+ )
887
+
888
+ # 4. Prepare latent variables
889
+ num_channels_latents = self.transformer.config.in_channels // 4
890
+ if do_rf_inversion:
891
+ latents = inverted_latents
892
+ else:
893
+ latents, latent_image_ids = self.prepare_latents(
894
+ batch_size * num_images_per_prompt,
895
+ num_channels_latents,
896
+ height,
897
+ width,
898
+ prompt_embeds.dtype,
899
+ device,
900
+ generator,
901
+ latents,
902
+ )
903
+
904
+ # 5. Prepare timesteps
905
+ sigmas = (
906
+ np.linspace(1.0, 1 / num_inference_steps, num_inference_steps)
907
+ if sigmas is None
908
+ else sigmas
909
+ )
910
+ image_seq_len = (int(height) // self.vae_scale_factor // 2) * (
911
+ int(width) // self.vae_scale_factor // 2
912
+ )
913
+ mu = calculate_shift(
914
+ image_seq_len,
915
+ self.scheduler.config.base_image_seq_len,
916
+ self.scheduler.config.max_image_seq_len,
917
+ self.scheduler.config.base_shift,
918
+ self.scheduler.config.max_shift,
919
+ )
920
+ timesteps, num_inference_steps = retrieve_timesteps(
921
+ self.scheduler,
922
+ num_inference_steps,
923
+ device,
924
+ timesteps,
925
+ sigmas,
926
+ mu=mu,
927
+ )
928
+ if do_rf_inversion:
929
+ start_timestep = int(start_timestep * num_inference_steps)
930
+ stop_timestep = min(
931
+ int(stop_timestep * num_inference_steps), num_inference_steps
932
+ )
933
+ timesteps, sigmas, num_inference_steps = self.get_timesteps(
934
+ num_inference_steps, strength
935
+ )
936
+ num_warmup_steps = max(
937
+ len(timesteps) - num_inference_steps * self.scheduler.order, 0
938
+ )
939
+ self._num_timesteps = len(timesteps)
940
+
941
+ # handle guidance
942
+ if self.transformer.config.guidance_embeds:
943
+ guidance = torch.full(
944
+ [1], guidance_scale, device=device, dtype=torch.float32
945
+ )
946
+ guidance = guidance.expand(latents.shape[0])
947
+ else:
948
+ guidance = None
949
+
950
+ if do_rf_inversion:
951
+ y_0 = image_latents.clone()
952
+ # 6. Denoising loop / Controlled Reverse ODE, Algorithm 2 from: https://arxiv.org/pdf/2410.10792
953
+ with self.progress_bar(total=num_inference_steps) as progress_bar:
954
+ for i, t in enumerate(timesteps):
955
+ if do_rf_inversion:
956
+ # ti (current timestep) as annotated in algorithm 2 - i/num_inference_steps.
957
+ t_i = 1 - t / 1000
958
+ dt = torch.tensor(1 / (len(timesteps) - 1), device=device)
959
+
960
+ if self.interrupt:
961
+ continue
962
+
963
+ # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
964
+ timestep = t.expand(latents.shape[0]).to(latents.dtype)
965
+
966
+ noise_pred = self.transformer(
967
+ hidden_states=latents,
968
+ timestep=timestep / 1000,
969
+ guidance=guidance,
970
+ pooled_projections=pooled_prompt_embeds,
971
+ encoder_hidden_states=prompt_embeds,
972
+ txt_ids=text_ids,
973
+ img_ids=latent_image_ids,
974
+ joint_attention_kwargs=self.joint_attention_kwargs,
975
+ return_dict=False,
976
+ )[0]
977
+
978
+ latents_dtype = latents.dtype
979
+ if do_rf_inversion:
980
+ v_t = -noise_pred
981
+ v_t_cond = (y_0 - latents) / (1 - t_i)
982
+ eta_t = eta if start_timestep <= i < stop_timestep else 0.0
983
+ if decay_eta:
984
+ eta_t = (
985
+ eta_t * (1 - i / num_inference_steps) ** eta_decay_power
986
+ ) # Decay eta over the loop
987
+ v_hat_t = v_t + eta_t * (v_t_cond - v_t)
988
+
989
+ # SDE Eq: 17 from https://arxiv.org/pdf/2410.10792
990
+ latents = latents + v_hat_t * (sigmas[i] - sigmas[i + 1])
991
+ else:
992
+ # compute the previous noisy sample x_t -> x_t-1
993
+ latents = self.scheduler.step(
994
+ noise_pred, t, latents, return_dict=False
995
+ )[0]
996
+
997
+ if latents.dtype != latents_dtype:
998
+ if torch.backends.mps.is_available():
999
+ # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
1000
+ latents = latents.to(latents_dtype)
1001
+
1002
+ if callback_on_step_end is not None:
1003
+ callback_kwargs = {}
1004
+ for k in callback_on_step_end_tensor_inputs:
1005
+ callback_kwargs[k] = locals()[k]
1006
+ callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
1007
+
1008
+ latents = callback_outputs.pop("latents", latents)
1009
+ prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
1010
+
1011
+ # call the callback, if provided
1012
+ if i == len(timesteps) - 1 or (
1013
+ (i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0
1014
+ ):
1015
+ progress_bar.update()
1016
+
1017
+ if XLA_AVAILABLE:
1018
+ xm.mark_step()
1019
+
1020
+ if output_type == "latent":
1021
+ image = latents
1022
+
1023
+ else:
1024
+ latents = self._unpack_latents(
1025
+ latents, height, width, self.vae_scale_factor
1026
+ )
1027
+ latents = (
1028
+ latents / self.vae.config.scaling_factor
1029
+ ) + self.vae.config.shift_factor
1030
+ image = self.vae.decode(latents, return_dict=False)[0]
1031
+ image = self.image_processor.postprocess(image, output_type=output_type)
1032
+
1033
+ # Offload all models
1034
+ self.maybe_free_model_hooks()
1035
+
1036
+ if not return_dict:
1037
+ return (image,)
1038
+
1039
+ return FluxPipelineOutput(images=image)
1040
+
1041
+ @torch.no_grad()
1042
+ def invert(
1043
+ self,
1044
+ image: PipelineImageInput,
1045
+ source_prompt: str = "",
1046
+ source_guidance_scale=0.0,
1047
+ num_inversion_steps: int = 28,
1048
+ strength: float = 1.0,
1049
+ gamma: float = 0.5,
1050
+ height: Optional[int] = None,
1051
+ width: Optional[int] = None,
1052
+ timesteps: List[int] = None,
1053
+ dtype: Optional[torch.dtype] = None,
1054
+ joint_attention_kwargs: Optional[Dict[str, Any]] = None,
1055
+ ):
1056
+ r"""
1057
+ Performs Algorithm 1: Controlled Forward ODE from https://arxiv.org/pdf/2410.10792
1058
+ Args:
1059
+ image (`PipelineImageInput`):
1060
+ Input for the image(s) that are to be edited. Multiple input images have to default to the same aspect
1061
+ ratio.
1062
+ source_prompt (`str` or `List[str]`, *optional* defaults to an empty prompt as done in the original paper):
1063
+ The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
1064
+ instead.
1065
+ source_guidance_scale (`float`, *optional*, defaults to 0.0):
1066
+ Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
1067
+ `guidance_scale` is defined as `w` of equation 2. of [Imagen
1068
+ Paper](https://arxiv.org/pdf/2205.11487.pdf). For this algorithm, it's better to keep it 0.
1069
+ num_inversion_steps (`int`, *optional*, defaults to 28):
1070
+ The number of discretization steps.
1071
+ height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
1072
+ The height in pixels of the generated image. This is set to 1024 by default for the best results.
1073
+ width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
1074
+ The width in pixels of the generated image. This is set to 1024 by default for the best results.
1075
+ gamma (`float`, *optional*, defaults to 0.5):
1076
+ The controller guidance for the forward ODE, balancing faithfulness & editability:
1077
+ higher eta - better faithfullness, less editability. For more significant edits, lower the value of eta.
1078
+ timesteps (`List[int]`, *optional*):
1079
+ Custom timesteps to use for the denoising process with schedulers which support a `timesteps` argument
1080
+ in their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is
1081
+ passed will be used. Must be in descending order.
1082
+ """
1083
+ dtype = dtype or self.text_encoder.dtype
1084
+ batch_size = 1
1085
+ self._joint_attention_kwargs = joint_attention_kwargs
1086
+ num_channels_latents = self.transformer.config.in_channels // 4
1087
+
1088
+ height = height or self.default_sample_size * self.vae_scale_factor
1089
+ width = width or self.default_sample_size * self.vae_scale_factor
1090
+ device = self._execution_device
1091
+
1092
+ # 1. prepare image
1093
+ image_latents, _ = self.encode_image(
1094
+ image, height=height, width=width, dtype=dtype
1095
+ )
1096
+ image_latents, latent_image_ids = self.prepare_latents_inversion(
1097
+ batch_size,
1098
+ num_channels_latents,
1099
+ height,
1100
+ width,
1101
+ dtype,
1102
+ device,
1103
+ image_latents,
1104
+ )
1105
+
1106
+ # 2. prepare timesteps
1107
+ sigmas = np.linspace(1.0, 1 / num_inversion_steps, num_inversion_steps)
1108
+ image_seq_len = (int(height) // self.vae_scale_factor // 2) * (
1109
+ int(width) // self.vae_scale_factor // 2
1110
+ )
1111
+ mu = calculate_shift(
1112
+ image_seq_len,
1113
+ self.scheduler.config.base_image_seq_len,
1114
+ self.scheduler.config.max_image_seq_len,
1115
+ self.scheduler.config.base_shift,
1116
+ self.scheduler.config.max_shift,
1117
+ )
1118
+ timesteps, num_inversion_steps = retrieve_timesteps(
1119
+ self.scheduler,
1120
+ num_inversion_steps,
1121
+ device,
1122
+ timesteps,
1123
+ sigmas,
1124
+ mu=mu,
1125
+ )
1126
+ timesteps, sigmas, num_inversion_steps = self.get_timesteps(
1127
+ num_inversion_steps, strength
1128
+ )
1129
+
1130
+ # 3. prepare text embeddings
1131
+ (
1132
+ prompt_embeds,
1133
+ pooled_prompt_embeds,
1134
+ text_ids,
1135
+ ) = self.encode_prompt(
1136
+ prompt=source_prompt,
1137
+ prompt_2=source_prompt,
1138
+ device=device,
1139
+ )
1140
+ # 4. handle guidance
1141
+ if self.transformer.config.guidance_embeds:
1142
+ guidance = torch.full(
1143
+ [1], source_guidance_scale, device=device, dtype=torch.float32
1144
+ )
1145
+ else:
1146
+ guidance = None
1147
+
1148
+ # Eq 8 dY_t = [u_t(Y_t) + γ(u_t(Y_t|y_1) - u_t(Y_t))]dt
1149
+ Y_t = image_latents
1150
+ y_1 = torch.randn_like(Y_t)
1151
+ N = len(sigmas)
1152
+
1153
+ # forward ODE loop
1154
+ with self.progress_bar(total=N - 1) as progress_bar:
1155
+ for i in range(N - 1):
1156
+ t_i = torch.tensor(i / (N), dtype=Y_t.dtype, device=device)
1157
+ timestep = torch.tensor(t_i, dtype=Y_t.dtype, device=device).repeat(
1158
+ batch_size
1159
+ )
1160
+
1161
+ # get the unconditional vector field
1162
+ u_t_i = self.transformer(
1163
+ hidden_states=Y_t,
1164
+ timestep=timestep,
1165
+ guidance=guidance,
1166
+ pooled_projections=pooled_prompt_embeds,
1167
+ encoder_hidden_states=prompt_embeds,
1168
+ txt_ids=text_ids,
1169
+ img_ids=latent_image_ids,
1170
+ joint_attention_kwargs=self.joint_attention_kwargs,
1171
+ return_dict=False,
1172
+ )[0]
1173
+
1174
+ # get the conditional vector field
1175
+ u_t_i_cond = (y_1 - Y_t) / (1 - t_i)
1176
+
1177
+ # controlled vector field
1178
+ # Eq 8 dY_t = [u_t(Y_t) + γ(u_t(Y_t|y_1) - u_t(Y_t))]dt
1179
+ u_hat_t_i = u_t_i + gamma * (u_t_i_cond - u_t_i)
1180
+ Y_t = Y_t + u_hat_t_i * (sigmas[i] - sigmas[i + 1])
1181
+ progress_bar.update()
1182
+
1183
+ # return the inverted latents (start point for the denoising loop), encoded image & latent image ids
1184
+ return Y_t, image_latents, latent_image_ids