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,18 @@
1
+ import torch
2
+ from huggingface_hub import get_token
3
+ import requests
4
+ import io
5
+
6
+
7
+ def remote_text_encoder(prompts, url, device):
8
+ response = requests.post(
9
+ url,
10
+ json={"prompt": prompts},
11
+ headers={
12
+ "Authorization": f"Bearer {get_token()}",
13
+ "Content-Type": "application/json",
14
+ },
15
+ )
16
+ prompt_embeds = torch.load(io.BytesIO(response.content))
17
+
18
+ return prompt_embeds.to(device)
dw/previous_results.py ADDED
@@ -0,0 +1,259 @@
1
+ import logging
2
+ from itertools import product
3
+
4
+ from .arguments import (
5
+ FROM_PREVIOUS_RESULT_KEY,
6
+ PREVIOUS_RESULT_PREFIX,
7
+ build_objects,
8
+ )
9
+
10
+ logger = logging.getLogger("dw")
11
+
12
+ # Maximum number of iterations to prevent resource exhaustion
13
+ MAX_ITERATIONS = 10000
14
+
15
+
16
+ def get_iterations(argument_template, previous_results):
17
+ """Generate argument combinations using previous task results.
18
+
19
+ Takes a template of arguments and expands any references to previous results
20
+ into all possible combinations of those results.
21
+
22
+ Args:
23
+ argument_template: Dict or list containing argument definitions
24
+ previous_results: Dict of results from previously executed steps
25
+
26
+ Returns:
27
+ List of argument dictionaries, one for each possible combination
28
+ """
29
+ # Special case: if template is a list, use it directly without processing
30
+ if isinstance(argument_template, list):
31
+ logger.debug("Using list argument template directly")
32
+ return argument_template
33
+
34
+ # Find any references to previous results in the template
35
+ # Returns dict of {arg_key: result_reference}
36
+ result_refs = find_previous_result_refs(argument_template)
37
+
38
+ # If no references found, return the template as-is
39
+ if not result_refs:
40
+ logger.debug("No result references found in template")
41
+ # Shallow copy: realize_args may have already loaded large media
42
+ # (PIL images, full video frame lists) into the template, so a deep
43
+ # copy would multiply memory use. Contract: iteration dicts may only
44
+ # be mutated at the top level (key pop/assign); nested values are
45
+ # shared across iterations and must never be mutated in place.
46
+ return [dict(argument_template)]
47
+
48
+ logger.debug(f"Found {len(result_refs)} result references: {result_refs}")
49
+
50
+ # Create a dictionary mapping each reference path to its possible values
51
+ # Example: {('image',): [img1, img2], ('prompt',): ['text1', 'text2']}
52
+ ref_results = {
53
+ ref_path: list(get_previous_results(previous_results, ref_value))
54
+ for ref_path, ref_value in result_refs.items()
55
+ }
56
+
57
+ # Generate all possible combinations of argument values
58
+ keys = list(ref_results.keys())
59
+ iterations = []
60
+
61
+ # Use itertools.product to create cartesian product of all possible values
62
+ # Example: if ref_results has 2 images and 2 prompts, creates 4 combinations
63
+ for values in product(*[ref_results[k] for k in keys]):
64
+ # Create fresh shallow copy of template for each combination.
65
+ # Nested values (e.g. loaded PIL images, video frame lists) are
66
+ # shared across iterations, not deep-copied, to avoid multiplying
67
+ # media memory usage by the iteration count. Contract: iteration
68
+ # dicts may only be mutated at the top level (key pop/assign);
69
+ # nested values must never be mutated in place.
70
+ arguments = dict(argument_template)
71
+
72
+ # Replace each reference with its actual value
73
+ for path, value in zip(keys, values):
74
+ # Handle nested dictionary properties
75
+ # If value is dict and contains the key we're looking for, use that property
76
+ key = path[-1]
77
+ arguments = substitute_at_path(
78
+ arguments,
79
+ path,
80
+ value[key] if isinstance(value, dict) and key in value else value,
81
+ )
82
+
83
+ # Now that the media exists, build the objects that were waiting for it -
84
+ # a reference constructed from a step's output rather than from a file
85
+ iterations.append(build_objects(arguments))
86
+
87
+ # Safety check to prevent cartesian product explosion
88
+ if len(iterations) > MAX_ITERATIONS:
89
+ raise ValueError(
90
+ f"Too many iterations generated: {len(iterations)} exceeds maximum of {MAX_ITERATIONS}. "
91
+ f"This usually indicates too many previous_result references creating a cartesian product. "
92
+ f"Consider reducing the number of multi-value results or splitting into multiple steps."
93
+ )
94
+
95
+ logger.debug(f"Generated {len(iterations)} argument combinations")
96
+ return iterations
97
+
98
+
99
+ def get_previous_results(previous_results, previous_result_name):
100
+ """Retrieve results or specific properties from previous tasks.
101
+
102
+ Args:
103
+ previous_results: Dict of results from previous steps
104
+ previous_result_name: String identifying the result, optionally with property
105
+ Format: "step_name" or "step_name.property_name"
106
+
107
+ Returns:
108
+ List of results or specific properties from the referenced step
109
+ """
110
+ # Step names are unrestricted strings and may themselves contain dots
111
+ # (e.g. "v1.0"), so resolve against the known step names rather than
112
+ # blindly splitting on the first/only ".".
113
+
114
+ # Exact match: the whole reference is a known step name, no property.
115
+ if previous_result_name in previous_results:
116
+ logger.debug(f"Getting all artifacts from result {previous_result_name}")
117
+ return previous_results[previous_result_name].get_artifacts()
118
+
119
+ if "." not in previous_result_name:
120
+ raise KeyError(
121
+ f"Previous result '{previous_result_name}' not found. Available results: {list(previous_results.keys())}"
122
+ )
123
+
124
+ # Find the longest known step name that is a prefix of the reference
125
+ # followed by ".", and treat the remainder as the property name.
126
+ result_name = max(
127
+ (
128
+ name
129
+ for name in previous_results
130
+ if previous_result_name.startswith(name + ".")
131
+ ),
132
+ key=len,
133
+ default=None,
134
+ )
135
+
136
+ if result_name is None:
137
+ raise KeyError(
138
+ f"Previous result '{previous_result_name}' not found. Available results: {list(previous_results.keys())}"
139
+ )
140
+
141
+ property_name = previous_result_name[len(result_name) + 1 :]
142
+ logger.debug(f"Getting property {property_name} from result {result_name}")
143
+ return previous_results[result_name].get_artifact_properties(property_name)
144
+
145
+
146
+ def resolve_chain_prompts(step_action, previous_results):
147
+ """Resolve a pipeline chain's per-segment prompts against previous results.
148
+
149
+ A chain's "prompts" list is not part of the step's argument template, so the
150
+ cartesian pass that expands "previous_result:" everywhere else never reaches
151
+ it. That matters for a chain whose opening segment is written by a different
152
+ step from the ones that continue it - a continuation prompt declares a video
153
+ reference the first segment does not have.
154
+
155
+ Each entry resolves independently and yields one prompt, so this never
156
+ multiplies iterations the way an argument reference does; a reference that
157
+ produced several artifacts uses the first.
158
+
159
+ The resolved list is left on the step action for run_chain to pick up, and
160
+ nothing happens at all for a pipeline without chain prompts.
161
+ """
162
+ definition = getattr(step_action, "pipeline_definition", None)
163
+ if not isinstance(definition, dict):
164
+ return
165
+
166
+ chain = definition.get("chain", None) or {}
167
+ prompts = chain.get("prompts", None)
168
+ if not prompts:
169
+ return
170
+
171
+ resolved = []
172
+ for entry in prompts:
173
+ if isinstance(entry, str) and entry.startswith("previous_result:"):
174
+ artifacts = get_previous_results(
175
+ previous_results, entry.removeprefix("previous_result:")
176
+ )
177
+ if not artifacts:
178
+ raise ValueError(f"Chain prompt reference '{entry}' produced no result")
179
+ if len(artifacts) > 1:
180
+ logger.warning(
181
+ f"Chain prompt reference '{entry}' produced {len(artifacts)} "
182
+ f"results - using the first"
183
+ )
184
+ entry = artifacts[0]
185
+ resolved.append(entry)
186
+
187
+ step_action.chain_prompts = resolved
188
+
189
+
190
+ def find_previous_result_refs(arguments):
191
+ """Find all values in an argument structure that reference previous results.
192
+
193
+ A reference is written either as a value with the "previous_result:" prefix, or
194
+ as the step name a 'from_previous_result' object description is built from. Both
195
+ are found at any depth: an argument that takes a constructed object holds it
196
+ inside a list - MiniMax-H3's 'references' - so the reference is nested rather
197
+ than sitting at the top of the arguments.
198
+
199
+ Args:
200
+ arguments: Dictionary of argument definitions
201
+
202
+ Returns:
203
+ Dict mapping the path of each reference to the result name it names. A path
204
+ is the tuple of keys and list indices that reaches the value, so a top-level
205
+ {'image': 'previous_result:step1'} comes back as {('image',): 'step1'}
206
+ """
207
+ found = {}
208
+ _collect_refs(arguments, (), found)
209
+ return found
210
+
211
+
212
+ def _collect_refs(value, path, found):
213
+ """Walk an argument structure, collecting every reference by its path."""
214
+ if isinstance(value, dict):
215
+ for key, item in value.items():
216
+ # The object description names its step bare, the way it would name a
217
+ # file - the prefix would only repeat what the key already says
218
+ if key == FROM_PREVIOUS_RESULT_KEY and isinstance(item, str):
219
+ found[path + (key,)] = item
220
+ else:
221
+ _collect_refs(item, path + (key,), found)
222
+
223
+ elif isinstance(value, list):
224
+ for index, item in enumerate(value):
225
+ _collect_refs(item, path + (index,), found)
226
+
227
+ elif isinstance(value, str) and value.startswith(PREVIOUS_RESULT_PREFIX):
228
+ found[path] = value[len(PREVIOUS_RESULT_PREFIX) :]
229
+
230
+
231
+ def substitute_at_path(container, path, value):
232
+ """A copy of container with value placed at path.
233
+
234
+ Only the containers along the path are copied. Everything beside them stays
235
+ shared, which is the same contract the top-level copy keeps: iterations share
236
+ their nested values, so a substitution deep in one must not be visible in the
237
+ others.
238
+
239
+ Args:
240
+ container: The dict or list to substitute into
241
+ path: Tuple of keys and indices reaching the value to replace
242
+ value: What to put there
243
+
244
+ Returns:
245
+ The copied container
246
+ """
247
+ key = path[0]
248
+ replacement = (
249
+ value if len(path) == 1 else substitute_at_path(container[key], path[1:], value)
250
+ )
251
+
252
+ if isinstance(container, list):
253
+ copied = list(container)
254
+ copied[key] = replacement
255
+ return copied
256
+
257
+ copied = dict(container)
258
+ copied[key] = replacement
259
+ return copied
dw/prompt_weighting.py ADDED
@@ -0,0 +1,378 @@
1
+ """
2
+ Prompt weighting and long prompt support for diffusers pipelines.
3
+
4
+ Parses A1111-style prompt syntax: (word:1.5) for emphasis, [word] for de-emphasis,
5
+ ((word)) for nested weighting. Supports prompts longer than the 77-token CLIP limit.
6
+
7
+ Produces prompt_embeds tensors that replace the prompt string argument in pipeline calls.
8
+
9
+ Based on sd_embed by Andrew Zhu (https://github.com/xhinker/sd_embed)
10
+ License: Apache 2.0
11
+ """
12
+
13
+ import re
14
+ import gc
15
+ import logging
16
+ from typing import Tuple
17
+
18
+ import torch
19
+ from transformers import CLIPTokenizer, T5Tokenizer
20
+
21
+ logger = logging.getLogger("dw")
22
+
23
+
24
+ # ---------------------------------------------------------------------------
25
+ # Prompt parser — A1111-style (word:weight) syntax
26
+ # ---------------------------------------------------------------------------
27
+
28
+ _re_attention = re.compile(
29
+ r"""
30
+ \\\(|
31
+ \\\)|
32
+ \\\[|
33
+ \\]|
34
+ \\\\|
35
+ \\|
36
+ \(|
37
+ \[|
38
+ :\s*([+-]?[.\d]+)\s*\)|
39
+ \)|
40
+ ]|
41
+ [^\\()\[\]:]+|
42
+ :
43
+ """,
44
+ re.X,
45
+ )
46
+
47
+ _re_break = re.compile(r"\s*\bBREAK\b\s*", re.S)
48
+
49
+
50
+ def parse_prompt_attention(text):
51
+ """Parse a prompt string with attention weights.
52
+
53
+ Syntax:
54
+ (abc) — weight 1.1
55
+ (abc:1.5) — weight 1.5
56
+ ((abc)) — weight 1.21 (1.1 * 1.1)
57
+ [abc] — weight 1/1.1 ≈ 0.91
58
+ \\( \\) — literal parens
59
+
60
+ Returns list of [text, weight] pairs.
61
+ """
62
+ res = []
63
+ round_brackets = []
64
+ square_brackets = []
65
+ round_bracket_multiplier = 1.1
66
+ square_bracket_multiplier = 1 / 1.1
67
+
68
+ def multiply_range(start_position, multiplier):
69
+ for p in range(start_position, len(res)):
70
+ res[p][1] *= multiplier
71
+
72
+ for m in _re_attention.finditer(text):
73
+ text_match = m.group(0)
74
+ weight = m.group(1)
75
+
76
+ if text_match.startswith("\\"):
77
+ res.append([text_match[1:], 1.0])
78
+ elif text_match == "(":
79
+ round_brackets.append(len(res))
80
+ elif text_match == "[":
81
+ square_brackets.append(len(res))
82
+ elif weight is not None and len(round_brackets) > 0:
83
+ multiply_range(round_brackets.pop(), float(weight))
84
+ elif text_match == ")" and len(round_brackets) > 0:
85
+ multiply_range(round_brackets.pop(), round_bracket_multiplier)
86
+ elif text_match == "]" and len(square_brackets) > 0:
87
+ multiply_range(square_brackets.pop(), square_bracket_multiplier)
88
+ else:
89
+ parts = re.split(_re_break, text_match)
90
+ for i, part in enumerate(parts):
91
+ if i > 0:
92
+ res.append(["BREAK", -1])
93
+ res.append([part, 1.0])
94
+
95
+ for pos in round_brackets:
96
+ multiply_range(pos, round_bracket_multiplier)
97
+ for pos in square_brackets:
98
+ multiply_range(pos, square_bracket_multiplier)
99
+
100
+ if len(res) == 0:
101
+ res = [["", 1.0]]
102
+
103
+ # merge runs of identical weights
104
+ i = 0
105
+ while i + 1 < len(res):
106
+ if res[i][1] == res[i + 1][1]:
107
+ res[i][0] += res[i + 1][0]
108
+ res.pop(i + 1)
109
+ else:
110
+ i += 1
111
+
112
+ return res
113
+
114
+
115
+ # ---------------------------------------------------------------------------
116
+ # Tokenization helpers
117
+ # ---------------------------------------------------------------------------
118
+
119
+
120
+ def _tokenize_clip_with_weights(clip_tokenizer: CLIPTokenizer, prompt: str):
121
+ """Tokenize with CLIP and return (token_ids, weights)."""
122
+ if not prompt:
123
+ prompt = "empty"
124
+
125
+ texts_and_weights = parse_prompt_attention(prompt)
126
+ text_tokens, text_weights = [], []
127
+ for word, weight in texts_and_weights:
128
+ token = clip_tokenizer(word, truncation=False).input_ids[1:-1]
129
+ text_tokens.extend(token)
130
+ text_weights.extend([weight] * len(token))
131
+ return text_tokens, text_weights
132
+
133
+
134
+ def _tokenize_t5_with_weights(t5_tokenizer: T5Tokenizer, prompt: str):
135
+ """Tokenize with T5 and return (token_ids, weights)."""
136
+ if not prompt:
137
+ prompt = "empty"
138
+
139
+ texts_and_weights = parse_prompt_attention(prompt)
140
+ text_tokens, text_weights = [], []
141
+ for word, weight in texts_and_weights:
142
+ token = t5_tokenizer(word, truncation=False, add_special_tokens=True).input_ids
143
+ text_tokens.extend(token)
144
+ text_weights.extend([weight] * len(token))
145
+ return text_tokens, text_weights
146
+
147
+
148
+ def _group_tokens_and_weights(token_ids, weights, pad_last_block=True):
149
+ """Group tokens into 77-token chunks with BOS/EOS padding."""
150
+ bos, eos = 49406, 49407
151
+
152
+ new_token_ids = []
153
+ new_weights = []
154
+
155
+ # work on copies to avoid mutating originals
156
+ token_ids = list(token_ids)
157
+ weights = list(weights)
158
+
159
+ while len(token_ids) >= 75:
160
+ head_tokens = [token_ids.pop(0) for _ in range(75)]
161
+ head_weights = [weights.pop(0) for _ in range(75)]
162
+ new_token_ids.append([bos] + head_tokens + [eos])
163
+ new_weights.append([1.0] + head_weights + [1.0])
164
+
165
+ if len(token_ids) > 0:
166
+ padding_len = 75 - len(token_ids) if pad_last_block else 0
167
+ new_token_ids.append([bos] + token_ids + [eos] * padding_len + [eos])
168
+ new_weights.append([1.0] + weights + [1.0] * padding_len + [1.0])
169
+
170
+ return new_token_ids, new_weights
171
+
172
+
173
+ # ---------------------------------------------------------------------------
174
+ # Pipeline-specific weighted embedding functions
175
+ # ---------------------------------------------------------------------------
176
+
177
+
178
+ def _get_device(pipeline):
179
+ """Get the appropriate compute device for the pipeline."""
180
+ device = pipeline.device
181
+ # Offloaded pipelines report cpu (model offload) or meta (sequential offload)
182
+ if device is not None and device.type not in ("cpu", "meta"):
183
+ return device
184
+
185
+ # Fall back to the device dw is running on, which honors the DW_DEVICE and
186
+ # settings overrides - never a hardcoded accelerator
187
+ from . import get_device
188
+
189
+ return torch.device(get_device())
190
+
191
+
192
+ def _hook_managed(module):
193
+ """Whether accelerate offload hooks own this module's placement.
194
+
195
+ Moving a hooked module manually fights the hooks - and raises outright on the
196
+ meta tensors sequential offload leaves behind. The hooks bring the module to
197
+ its execution device on forward, so no manual move is needed.
198
+ """
199
+ return getattr(module, "_hf_hook", None) is not None
200
+
201
+
202
+ def get_weighted_text_embeddings_flux(
203
+ pipe,
204
+ prompt: str = "",
205
+ prompt2: str = None,
206
+ device=None,
207
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
208
+ """Generate weighted text embeddings for Flux pipelines.
209
+
210
+ Supports long prompts (beyond 77 tokens) and A1111-style weighting syntax.
211
+
212
+ Args:
213
+ pipe: A loaded FluxPipeline with tokenizer, tokenizer_2, text_encoder, text_encoder_2
214
+ prompt: Primary prompt with optional weighting syntax
215
+ prompt2: Optional second prompt for T5 encoder (defaults to prompt)
216
+ device: Target device override
217
+
218
+ Returns:
219
+ (prompt_embeds, pooled_prompt_embeds) — pass directly to pipe() as kwargs
220
+ """
221
+ prompt2 = prompt if prompt2 is None else prompt2
222
+
223
+ target_device = device if device is not None else _get_device(pipe)
224
+
225
+ # Move text encoders to device if the pipeline sits on the CPU - unless
226
+ # offload hooks manage them, in which case they move themselves on forward
227
+ encoders_moved = False
228
+ if pipe.device.type == "cpu" and not (
229
+ _hook_managed(pipe.text_encoder) or _hook_managed(pipe.text_encoder_2)
230
+ ):
231
+ pipe.text_encoder.to(target_device)
232
+ pipe.text_encoder_2.to(target_device)
233
+ encoders_moved = True
234
+
235
+ # Tokenize with CLIP (tokenizer 1) for pooled embeddings
236
+ prompt_tokens, prompt_weights = _tokenize_clip_with_weights(pipe.tokenizer, prompt)
237
+ prompt_token_groups, _ = _group_tokens_and_weights(prompt_tokens, prompt_weights)
238
+
239
+ # Generate pooled CLIP embeddings (mean across token groups)
240
+ pool_embeds_list = []
241
+ for token_group in prompt_token_groups:
242
+ token_tensor = torch.tensor(
243
+ [token_group], dtype=torch.long, device=target_device
244
+ )
245
+ with torch.no_grad():
246
+ embeds = pipe.text_encoder(token_tensor, output_hidden_states=False)
247
+ pool_embeds_list.append(embeds.pooler_output.squeeze(0))
248
+
249
+ pooled_prompt_embeds = torch.stack(pool_embeds_list, dim=0)
250
+ pooled_prompt_embeds = pooled_prompt_embeds.mean(dim=0, keepdim=True)
251
+ pooled_prompt_embeds = pooled_prompt_embeds.to(
252
+ dtype=pipe.text_encoder.dtype, device=target_device
253
+ )
254
+
255
+ # Tokenize with T5 (tokenizer 2) for main prompt embeddings
256
+ prompt_tokens_2, prompt_weights_2 = _tokenize_t5_with_weights(
257
+ pipe.tokenizer_2, prompt2
258
+ )
259
+
260
+ token_tensor_2 = torch.tensor([prompt_tokens_2], dtype=torch.long)
261
+ with torch.no_grad():
262
+ t5_embeds = pipe.text_encoder_2(token_tensor_2.to(target_device))[0].squeeze(0)
263
+ t5_embeds = t5_embeds.to(device=target_device)
264
+
265
+ # Apply per-token weights to T5 embeddings
266
+ for i in range(len(prompt_weights_2)):
267
+ if prompt_weights_2[i] != 1.0:
268
+ t5_embeds[i] = t5_embeds[i] * prompt_weights_2[i]
269
+
270
+ prompt_embeds = t5_embeds.unsqueeze(0).to(
271
+ dtype=pipe.text_encoder_2.dtype, device=target_device
272
+ )
273
+
274
+ # Release encoders back to CPU if we moved them
275
+ if encoders_moved:
276
+ pipe.text_encoder.to("cpu")
277
+ pipe.text_encoder_2.to("cpu")
278
+ gc.collect()
279
+
280
+ from . import empty_device_cache
281
+
282
+ empty_device_cache()
283
+
284
+ return prompt_embeds, pooled_prompt_embeds
285
+
286
+
287
+ # ---------------------------------------------------------------------------
288
+ # Dispatcher — selects the right function based on pipeline type
289
+ # ---------------------------------------------------------------------------
290
+
291
+ # Map pipeline class names to their embedding functions
292
+ _PIPELINE_FUNCTIONS = {
293
+ "FluxPipeline": get_weighted_text_embeddings_flux,
294
+ "FluxImg2ImgPipeline": get_weighted_text_embeddings_flux,
295
+ "FluxInpaintPipeline": get_weighted_text_embeddings_flux,
296
+ "FluxControlNetPipeline": get_weighted_text_embeddings_flux,
297
+ }
298
+
299
+ # The encoder stack get_weighted_text_embeddings_flux drives - CLIP for pooled
300
+ # embeddings, T5 for the main ones
301
+ _FLUX_ENCODER_STACK = ("tokenizer", "tokenizer_2", "text_encoder", "text_encoder_2")
302
+
303
+
304
+ def _select_embedding_function(pipeline):
305
+ """The weighted-embedding function for a pipeline, or None.
306
+
307
+ Exact class names first; any other Flux-family pipeline (FluxKontext,
308
+ FluxFill, a user subclass) carrying the same CLIP+T5 encoder stack uses the
309
+ flux function too, so support does not lag every new variant diffusers adds.
310
+ """
311
+ embed_fn = _PIPELINE_FUNCTIONS.get(pipeline.__class__.__name__)
312
+ if embed_fn is not None:
313
+ return embed_fn
314
+
315
+ if pipeline.__class__.__name__.startswith("Flux") and all(
316
+ getattr(pipeline, name, None) is not None for name in _FLUX_ENCODER_STACK
317
+ ):
318
+ logger.info(
319
+ f"{pipeline.__class__.__name__} carries the Flux encoder stack - "
320
+ f"applying flux prompt weighting"
321
+ )
322
+ return get_weighted_text_embeddings_flux
323
+
324
+ return None
325
+
326
+
327
+ def apply_prompt_weighting(pipeline, arguments, device=None):
328
+ """Apply prompt weighting to pipeline arguments if the prompt contains weight syntax.
329
+
330
+ Checks if the prompt uses weighting syntax. If so, generates weighted embeddings
331
+ and replaces the prompt string with embedding tensors in the arguments dict.
332
+
333
+ Args:
334
+ pipeline: The loaded diffusers pipeline
335
+ arguments: Mutable dict of pipeline call arguments
336
+ device: The device the pipeline runs on - embeddings are created there.
337
+ Defaults to the pipeline's own device, or the dw device when offloading
338
+ parks the pipeline on the CPU
339
+
340
+ Returns:
341
+ True if weighting was applied, False if prompt was left as-is.
342
+ """
343
+ prompt = arguments.get("prompt", None)
344
+ if prompt is None or not isinstance(prompt, str):
345
+ return False
346
+
347
+ # Quick check: does the prompt contain any weighting syntax?
348
+ if "(" not in prompt and "[" not in prompt:
349
+ return False
350
+
351
+ class_name = pipeline.__class__.__name__
352
+ embed_fn = _select_embedding_function(pipeline)
353
+ if embed_fn is None:
354
+ logger.warning(
355
+ f"Prompt weighting not supported for {class_name}. "
356
+ f"Supported: {', '.join(_PIPELINE_FUNCTIONS.keys())} and Flux-family "
357
+ f"pipelines with the CLIP+T5 encoder stack. "
358
+ f"Passing prompt as plain text."
359
+ )
360
+ return False
361
+
362
+ logger.info(f"Applying prompt weighting for {class_name}")
363
+ prompt2 = arguments.pop("prompt_2", None)
364
+ prompt_str = arguments.pop("prompt")
365
+
366
+ prompt_embeds, pooled_prompt_embeds = embed_fn(
367
+ pipeline, prompt=prompt_str, prompt2=prompt2, device=device
368
+ )
369
+
370
+ arguments["prompt_embeds"] = prompt_embeds
371
+ arguments["pooled_prompt_embeds"] = pooled_prompt_embeds
372
+
373
+ # Remove negative_prompt if present — can't mix string and embeds
374
+ if "negative_prompt" in arguments:
375
+ logger.debug("Removing negative_prompt (incompatible with prompt_embeds)")
376
+ arguments.pop("negative_prompt")
377
+
378
+ return True