diffusers-workflow 0.4.0__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 (260) hide show
  1. diffusers_workflow-0.4.0.dist-info/METADATA +318 -0
  2. diffusers_workflow-0.4.0.dist-info/RECORD +260 -0
  3. diffusers_workflow-0.4.0.dist-info/WHEEL +5 -0
  4. diffusers_workflow-0.4.0.dist-info/entry_points.txt +7 -0
  5. diffusers_workflow-0.4.0.dist-info/licenses/LICENSE +201 -0
  6. diffusers_workflow-0.4.0.dist-info/top_level.txt +2 -0
  7. dw/__init__.py +440 -0
  8. dw/adapter_compatibility.py +226 -0
  9. dw/arguments.py +1231 -0
  10. dw/assessment_rules.py +159 -0
  11. dw/assets.py +130 -0
  12. dw/cache_blocks.json +16 -0
  13. dw/cache_blocks.py +146 -0
  14. dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
  15. dw/content_types.py +150 -0
  16. dw/dissolve_frame_errors.py +121 -0
  17. dw/docs/ACCELERATION.md +352 -0
  18. dw/docs/AGENT_LOOP.md +95 -0
  19. dw/docs/DEPENDENCIES.md +91 -0
  20. dw/docs/IP_ADAPTER.md +109 -0
  21. dw/docs/LORAS.md +131 -0
  22. dw/docs/MCP.md +517 -0
  23. dw/docs/PROMPT_WEIGHTING.md +78 -0
  24. dw/docs/QUANTIZATION.md +230 -0
  25. dw/docs/RECIPES_24GB.md +201 -0
  26. dw/docs/RELEASING.md +195 -0
  27. dw/docs/REMOTE.md +140 -0
  28. dw/docs/REPL_COMMANDS.md +121 -0
  29. dw/docs/REPL_WORKER_GUIDE.md +51 -0
  30. dw/docs/SECURITY.md +272 -0
  31. dw/docs/SECURITY_QUICKREF.md +112 -0
  32. dw/docs/SERVER.md +679 -0
  33. dw/docs/TASKS.md +1741 -0
  34. dw/docs/TESTING.md +71 -0
  35. dw/docs/WORKFLOW_GUIDE.md +2038 -0
  36. dw/docs/WORKSPACES.md +316 -0
  37. dw/download_watch.py +335 -0
  38. dw/elision.py +306 -0
  39. dw/events.py +275 -0
  40. dw/for_each.py +409 -0
  41. dw/host_memory.py +258 -0
  42. dw/host_memory_projection.py +230 -0
  43. dw/hub_cache.py +432 -0
  44. dw/introspection.py +1228 -0
  45. dw/kernel_availability.py +208 -0
  46. dw/locations.py +599 -0
  47. dw/log_setup.py +45 -0
  48. dw/loudness.py +82 -0
  49. dw/media_audio.py +217 -0
  50. dw/media_frames.py +367 -0
  51. dw/media_info.py +297 -0
  52. dw/pipeline_processors/chain.py +821 -0
  53. dw/pipeline_processors/config_objects.py +237 -0
  54. dw/pipeline_processors/pipeline.py +2297 -0
  55. dw/pipeline_processors/remote.py +46 -0
  56. dw/plan.py +920 -0
  57. dw/previous_results.py +411 -0
  58. dw/probe_paths.py +59 -0
  59. dw/prompt_schema.json +48 -0
  60. dw/prompt_weighting.py +378 -0
  61. dw/prompts.py +159 -0
  62. dw/realize.py +250 -0
  63. dw/reference_limits.py +215 -0
  64. dw/reference_names.py +125 -0
  65. dw/repl.py +338 -0
  66. dw/repl_commands.py +836 -0
  67. dw/repl_worker.py +159 -0
  68. dw/result.py +1720 -0
  69. dw/result_fps.py +82 -0
  70. dw/run.py +162 -0
  71. dw/runs.py +768 -0
  72. dw/scalar_result_validation.py +97 -0
  73. dw/schema.py +283 -0
  74. dw/security.py +1038 -0
  75. dw/select_validation.py +115 -0
  76. dw/serve.py +277 -0
  77. dw/server/__init__.py +2 -0
  78. dw/server/app.py +4586 -0
  79. dw/server/assess.py +132 -0
  80. dw/server/catalog_shape.py +487 -0
  81. dw/server/enhancers.py +129 -0
  82. dw/server/exports.py +480 -0
  83. dw/server/guides.py +257 -0
  84. dw/server/jobs.py +1561 -0
  85. dw/server/mcp_mount.py +95 -0
  86. dw/server/netinfo.py +124 -0
  87. dw/server/observed_cost.py +379 -0
  88. dw/server/sysinfo.py +71 -0
  89. dw/server/ui/assets/abap-08VXUWAP.js +1 -0
  90. dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
  91. dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
  92. dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
  93. dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
  94. dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
  95. dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
  96. dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
  97. dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
  98. dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
  99. dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
  100. dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
  101. dw/server/ui/assets/css-DIMkf-bt.js +3 -0
  102. dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
  103. dw/server/ui/assets/cssMode-CPznxfY8.js +1 -0
  104. dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
  105. dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
  106. dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
  107. dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
  108. dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
  109. dw/server/ui/assets/editor.api-CpWcotrd.js +847 -0
  110. dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
  111. dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
  112. dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
  113. dw/server/ui/assets/freemarker2-CXtRM8N4.js +3 -0
  114. dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
  115. dw/server/ui/assets/go-C-y9NEjX.js +1 -0
  116. dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
  117. dw/server/ui/assets/handlebars-N7x-6NMY.js +1 -0
  118. dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
  119. dw/server/ui/assets/html-PhsdjHSr.js +1 -0
  120. dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
  121. dw/server/ui/assets/htmlMode-Dgj0SEok.js +1 -0
  122. dw/server/ui/assets/index-3Vw6WAPW.css +1 -0
  123. dw/server/ui/assets/index-DgrYhQd9.js +43 -0
  124. dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
  125. dw/server/ui/assets/java-BEtHBSE6.js +1 -0
  126. dw/server/ui/assets/javascript-BJqN9Qhv.js +1 -0
  127. dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
  128. dw/server/ui/assets/jsonMode-DbM4SWSv.js +7 -0
  129. dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
  130. dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
  131. dw/server/ui/assets/less-B9JPFI3C.js +2 -0
  132. dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
  133. dw/server/ui/assets/liquid-BWr8lEc4.js +1 -0
  134. dw/server/ui/assets/lspLanguageFeatures-C1iGuDyZ.js +4 -0
  135. dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
  136. dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
  137. dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
  138. dw/server/ui/assets/mdx-DAdMi_0p.js +1 -0
  139. dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
  140. dw/server/ui/assets/monaco--ixms01u.css +1 -0
  141. dw/server/ui/assets/monaco-BGCeEqaw.js +56 -0
  142. dw/server/ui/assets/msdax-DauUninz.js +1 -0
  143. dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
  144. dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
  145. dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
  146. dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
  147. dw/server/ui/assets/perl-oz_6vUea.js +1 -0
  148. dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
  149. dw/server/ui/assets/php-nr791fC2.js +1 -0
  150. dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
  151. dw/server/ui/assets/postiats-43DmfD33.js +1 -0
  152. dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
  153. dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
  154. dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
  155. dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
  156. dw/server/ui/assets/python-Bcn70HdC.js +1 -0
  157. dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
  158. dw/server/ui/assets/r-BwWrilGY.js +1 -0
  159. dw/server/ui/assets/razor-D1HmNnby.js +1 -0
  160. dw/server/ui/assets/redis-ClamHrr6.js +1 -0
  161. dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
  162. dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
  163. dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
  164. dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
  165. dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
  166. dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
  167. dw/server/ui/assets/scheme-BeGwcela.js +1 -0
  168. dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
  169. dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
  170. dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
  171. dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
  172. dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
  173. dw/server/ui/assets/sql-NEE52Syq.js +1 -0
  174. dw/server/ui/assets/st-DbInun42.js +1 -0
  175. dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
  176. dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
  177. dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
  178. dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
  179. dw/server/ui/assets/tsMode-D6u0XmOW.js +11 -0
  180. dw/server/ui/assets/twig-De2hgUGE.js +1 -0
  181. dw/server/ui/assets/typescript-BU6v-LMV.js +1 -0
  182. dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
  183. dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
  184. dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
  185. dw/server/ui/assets/workers-Cn7cTUKr.js +1 -0
  186. dw/server/ui/assets/xml--0LP2Lwk.js +1 -0
  187. dw/server/ui/assets/yaml-mpBg9jnt.js +1 -0
  188. dw/server/ui/index.html +17 -0
  189. dw/server/updater.py +192 -0
  190. dw/settings.py +98 -0
  191. dw/shot_span_preflight.py +116 -0
  192. dw/shots.py +359 -0
  193. dw/slice_preflight.py +148 -0
  194. dw/step.py +187 -0
  195. dw/step_cache.py +442 -0
  196. dw/subfolders.py +107 -0
  197. dw/task_domains.py +307 -0
  198. dw/tasks/assess.py +826 -0
  199. dw/tasks/audio_transcription.py +88 -0
  200. dw/tasks/audio_utils.py +1862 -0
  201. dw/tasks/background_remover.py +43 -0
  202. dw/tasks/borders.py +113 -0
  203. dw/tasks/compose_text.py +74 -0
  204. dw/tasks/concat_videos.py +300 -0
  205. dw/tasks/depth_estimator.py +54 -0
  206. dw/tasks/diffusion_upscale.py +109 -0
  207. dw/tasks/dissolve_videos.py +342 -0
  208. dw/tasks/format_messages.py +24 -0
  209. dw/tasks/gather.py +173 -0
  210. dw/tasks/grade.py +97 -0
  211. dw/tasks/image_to_text.py +43 -0
  212. dw/tasks/image_utils.py +764 -0
  213. dw/tasks/interpolate_frames.py +252 -0
  214. dw/tasks/judge.py +68 -0
  215. dw/tasks/model_cache.py +55 -0
  216. dw/tasks/pair_audio.py +268 -0
  217. dw/tasks/qr_code.py +19 -0
  218. dw/tasks/restore_faces.py +175 -0
  219. dw/tasks/rife_model.py +192 -0
  220. dw/tasks/segment.py +121 -0
  221. dw/tasks/select.py +111 -0
  222. dw/tasks/speech_generation.py +228 -0
  223. dw/tasks/stabilize.py +129 -0
  224. dw/tasks/task.py +920 -0
  225. dw/tasks/tensor_image.py +57 -0
  226. dw/tasks/text_generation.py +169 -0
  227. dw/tasks/text_sections.py +80 -0
  228. dw/tasks/upscale.py +203 -0
  229. dw/tasks/video_utils.py +624 -0
  230. dw/tasks/zoe_depth.py +71 -0
  231. dw/teacache.py +381 -0
  232. dw/teacache_models.json +99 -0
  233. dw/test.py +29 -0
  234. dw/type_helpers.py +231 -0
  235. dw/validate.py +68 -0
  236. dw/variable_constraints.py +444 -0
  237. dw/variables.py +443 -0
  238. dw/video_extensions.py +141 -0
  239. dw/vram_estimate.py +116 -0
  240. dw/worker.py +764 -0
  241. dw/workflow.py +2007 -0
  242. dw/workflow_schema.json +1346 -0
  243. dw/workflow_sources.py +383 -0
  244. dw/workflows/h3_context_ir.json +57 -0
  245. dw/workflows/test.json +31 -0
  246. dw/workspace.py +730 -0
  247. dw_mcp/__init__.py +6 -0
  248. dw_mcp/__main__.py +133 -0
  249. dw_mcp/assets.py +336 -0
  250. dw_mcp/authoring.py +114 -0
  251. dw_mcp/catalog.py +360 -0
  252. dw_mcp/client.py +486 -0
  253. dw_mcp/diagnose.py +371 -0
  254. dw_mcp/exports.py +84 -0
  255. dw_mcp/guides.py +35 -0
  256. dw_mcp/media.py +638 -0
  257. dw_mcp/models.py +97 -0
  258. dw_mcp/prompts.py +104 -0
  259. dw_mcp/server.py +1343 -0
  260. dw_mcp/workspaces.py +212 -0
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
dw/prompts.py ADDED
@@ -0,0 +1,159 @@
1
+ """The prompt library: stored prompts a workflow references by name.
2
+
3
+ A prompt is one JSON file under the prompt directory - its text plus the
4
+ metadata the library pages show (description, intended model, tags). A
5
+ workflow argument written as 'prompt:name' or 'prompt:folder/name' loads
6
+ the file's text at run time, so the prompt is shared by reference rather
7
+ than copied into every workflow that uses it.
8
+ """
9
+
10
+ import json
11
+ import logging
12
+ import os
13
+
14
+ from .schema import load_schema, validate_data
15
+ from .security import validate_prompt_path, validate_prompt_reference
16
+ from .workspace import PROMPTS_SUBDIR, discover_library, library_fallbacks
17
+
18
+ logger = logging.getLogger("dw")
19
+
20
+ # The prefix marking a value as a reference to a stored prompt. The name after it
21
+ # is rooted at the prompt directory, not the workflow file - prompts are a shared
22
+ # library, and the same reference means the same text from every workflow
23
+ PROMPT_PREFIX = "prompt:"
24
+
25
+ # The prefixes a stored prompt's text may not begin with. Resolved text is
26
+ # substituted where the reference stood, so text that itself looks like a
27
+ # reference would be resolved again - or worse, expand a step's iterations
28
+ RESERVED_TEXT_PREFIXES = (
29
+ "previous_result:",
30
+ "variable:",
31
+ "constant:",
32
+ "asset:",
33
+ "output:",
34
+ PROMPT_PREFIX,
35
+ )
36
+
37
+
38
+ def get_prompt_dir(base_dir=None):
39
+ """The directory stored prompts are rooted at.
40
+
41
+ DW_PROMPT_DIR names it explicitly - the server sets it from --prompt-dir,
42
+ and the spawned worker inherits it. Below that, see
43
+ workspace.discover_library for the shared precedence (a named workspace,
44
+ then ./prompts, then a walk up from base_dir, then the workspace's
45
+ prompts/ as the fallback).
46
+
47
+ Read at call time, not import time, so a test or worker sees the current
48
+ value.
49
+
50
+ Args:
51
+ base_dir: The workflow file's directory, when one anchors the search
52
+ """
53
+ return discover_library(PROMPTS_SUBDIR, "DW_PROMPT_DIR", base_dir)
54
+
55
+
56
+ def prompt_search_path(prompt_dir=None, base_dir=None):
57
+ """Every directory a 'prompt:' reference is looked for in, in order.
58
+
59
+ The library a save would write to comes first, then the read-only ones
60
+ an entry point put on the path (workspace.library_fallbacks - the
61
+ prompts a --examples-dir tree brings with it). A name found earlier
62
+ shadows the same name later, the way it does on the workflow search
63
+ path.
64
+
65
+ Args:
66
+ prompt_dir: The first directory; defaults to get_prompt_dir()
67
+ base_dir: The workflow file's directory, anchoring discovery when no
68
+ prompt directory is configured
69
+ """
70
+ primary = prompt_dir or get_prompt_dir(base_dir)
71
+ return [primary] + library_fallbacks(PROMPTS_SUBDIR, primary)
72
+
73
+
74
+ def resolve_prompt_reference(reference, prompt_dir=None, base_dir=None):
75
+ """Resolve a 'prompt:' reference to the file it names.
76
+
77
+ Args:
78
+ reference: The 'prompt:name' or 'prompt:folder/name' string
79
+ prompt_dir: Directory the name is rooted at; defaults to get_prompt_dir()
80
+ base_dir: The workflow file's directory, anchoring discovery when no
81
+ prompt directory is configured
82
+
83
+ Returns:
84
+ The validated absolute path of the prompt file
85
+
86
+ Raises:
87
+ InvalidInputError: If the name is not a valid prompt name
88
+ ValueError: If no prompt file exists under that name in any directory
89
+ on the search path
90
+ """
91
+ name = validate_prompt_reference(reference.removeprefix(PROMPT_PREFIX).strip())
92
+ roots = prompt_search_path(prompt_dir, base_dir)
93
+ for root in roots:
94
+ path = os.path.join(root, name + ".json")
95
+ if os.path.isfile(path):
96
+ return validate_prompt_path(path, root)
97
+ searched = ", ".join(roots)
98
+ raise ValueError(
99
+ f"No prompt named '{name}' in {searched} - a prompt reference names "
100
+ f"a .json file under the prompt directory, without the extension"
101
+ )
102
+
103
+
104
+ def load_prompt(path):
105
+ """Read and validate one prompt file.
106
+
107
+ Args:
108
+ path: Path of the prompt file, already validated
109
+
110
+ Returns:
111
+ The prompt as a dict
112
+
113
+ Raises:
114
+ ValueError: If the file is not JSON or does not match the prompt schema
115
+ """
116
+ try:
117
+ with open(path, "r", encoding="utf-8") as file:
118
+ data = json.load(file)
119
+ except json.JSONDecodeError as error:
120
+ raise ValueError(f"Prompt file {path} is not valid JSON: {error}") from error
121
+
122
+ status, message = validate_data(data, load_schema("prompt"))
123
+ if not status:
124
+ raise ValueError(f"Prompt file {path} is not a valid prompt: {message}")
125
+
126
+ return data
127
+
128
+
129
+ def fetch_prompt(reference, prompt_dir=None, base_dir=None):
130
+ """Read the text a 'prompt:' reference names.
131
+
132
+ Args:
133
+ reference: The 'prompt:name' or 'prompt:folder/name' string
134
+ prompt_dir: Directory the name is rooted at; defaults to get_prompt_dir()
135
+ base_dir: The workflow file's directory, anchoring discovery when no
136
+ prompt directory is configured
137
+
138
+ Returns:
139
+ The prompt file's text field
140
+
141
+ Raises:
142
+ ValueError: If the prompt is missing, invalid, or its text is itself
143
+ a reference
144
+ """
145
+ path = resolve_prompt_reference(reference, prompt_dir, base_dir)
146
+ text = load_prompt(path)["text"]
147
+
148
+ # Arguments are realized more than once, and iteration expansion scans the
149
+ # realized template - text that begins like a reference would be treated
150
+ # as one on the next pass, so it is data that may not masquerade as syntax
151
+ if text.startswith(RESERVED_TEXT_PREFIXES):
152
+ raise ValueError(
153
+ f"Prompt '{reference}' has text beginning with a reference prefix "
154
+ f"({', '.join(RESERVED_TEXT_PREFIXES)}) - a prompt's text may not "
155
+ f"itself be a reference"
156
+ )
157
+
158
+ logger.info(f"Loaded prompt {reference} from {path}")
159
+ return text