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.
- diffusers_workflow-0.4.0.dist-info/METADATA +318 -0
- diffusers_workflow-0.4.0.dist-info/RECORD +260 -0
- diffusers_workflow-0.4.0.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0.dist-info/entry_points.txt +7 -0
- diffusers_workflow-0.4.0.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0.dist-info/top_level.txt +2 -0
- dw/__init__.py +440 -0
- dw/adapter_compatibility.py +226 -0
- dw/arguments.py +1231 -0
- dw/assessment_rules.py +159 -0
- dw/assets.py +130 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +146 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/content_types.py +150 -0
- dw/dissolve_frame_errors.py +121 -0
- dw/docs/ACCELERATION.md +352 -0
- dw/docs/AGENT_LOOP.md +95 -0
- dw/docs/DEPENDENCIES.md +91 -0
- dw/docs/IP_ADAPTER.md +109 -0
- dw/docs/LORAS.md +131 -0
- dw/docs/MCP.md +517 -0
- dw/docs/PROMPT_WEIGHTING.md +78 -0
- dw/docs/QUANTIZATION.md +230 -0
- dw/docs/RECIPES_24GB.md +201 -0
- dw/docs/RELEASING.md +195 -0
- dw/docs/REMOTE.md +140 -0
- dw/docs/REPL_COMMANDS.md +121 -0
- dw/docs/REPL_WORKER_GUIDE.md +51 -0
- dw/docs/SECURITY.md +272 -0
- dw/docs/SECURITY_QUICKREF.md +112 -0
- dw/docs/SERVER.md +679 -0
- dw/docs/TASKS.md +1741 -0
- dw/docs/TESTING.md +71 -0
- dw/docs/WORKFLOW_GUIDE.md +2038 -0
- dw/docs/WORKSPACES.md +316 -0
- dw/download_watch.py +335 -0
- dw/elision.py +306 -0
- dw/events.py +275 -0
- dw/for_each.py +409 -0
- dw/host_memory.py +258 -0
- dw/host_memory_projection.py +230 -0
- dw/hub_cache.py +432 -0
- dw/introspection.py +1228 -0
- dw/kernel_availability.py +208 -0
- dw/locations.py +599 -0
- dw/log_setup.py +45 -0
- dw/loudness.py +82 -0
- dw/media_audio.py +217 -0
- dw/media_frames.py +367 -0
- dw/media_info.py +297 -0
- dw/pipeline_processors/chain.py +821 -0
- dw/pipeline_processors/config_objects.py +237 -0
- dw/pipeline_processors/pipeline.py +2297 -0
- dw/pipeline_processors/remote.py +46 -0
- dw/plan.py +920 -0
- dw/previous_results.py +411 -0
- dw/probe_paths.py +59 -0
- dw/prompt_schema.json +48 -0
- dw/prompt_weighting.py +378 -0
- dw/prompts.py +159 -0
- dw/realize.py +250 -0
- dw/reference_limits.py +215 -0
- dw/reference_names.py +125 -0
- dw/repl.py +338 -0
- dw/repl_commands.py +836 -0
- dw/repl_worker.py +159 -0
- dw/result.py +1720 -0
- dw/result_fps.py +82 -0
- dw/run.py +162 -0
- dw/runs.py +768 -0
- dw/scalar_result_validation.py +97 -0
- dw/schema.py +283 -0
- dw/security.py +1038 -0
- dw/select_validation.py +115 -0
- dw/serve.py +277 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +4586 -0
- dw/server/assess.py +132 -0
- dw/server/catalog_shape.py +487 -0
- dw/server/enhancers.py +129 -0
- dw/server/exports.py +480 -0
- dw/server/guides.py +257 -0
- dw/server/jobs.py +1561 -0
- dw/server/mcp_mount.py +95 -0
- dw/server/netinfo.py +124 -0
- dw/server/observed_cost.py +379 -0
- dw/server/sysinfo.py +71 -0
- dw/server/ui/assets/abap-08VXUWAP.js +1 -0
- dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
- dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
- dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
- dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
- dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
- dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
- dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
- dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
- dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
- dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
- dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
- dw/server/ui/assets/css-DIMkf-bt.js +3 -0
- dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
- dw/server/ui/assets/cssMode-CPznxfY8.js +1 -0
- dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
- dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
- dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
- dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
- dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
- dw/server/ui/assets/editor.api-CpWcotrd.js +847 -0
- dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
- dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
- dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
- dw/server/ui/assets/freemarker2-CXtRM8N4.js +3 -0
- dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
- dw/server/ui/assets/go-C-y9NEjX.js +1 -0
- dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
- dw/server/ui/assets/handlebars-N7x-6NMY.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-PhsdjHSr.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-Dgj0SEok.js +1 -0
- dw/server/ui/assets/index-3Vw6WAPW.css +1 -0
- dw/server/ui/assets/index-DgrYhQd9.js +43 -0
- dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
- dw/server/ui/assets/java-BEtHBSE6.js +1 -0
- dw/server/ui/assets/javascript-BJqN9Qhv.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-DbM4SWSv.js +7 -0
- dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
- dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
- dw/server/ui/assets/less-B9JPFI3C.js +2 -0
- dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
- dw/server/ui/assets/liquid-BWr8lEc4.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-C1iGuDyZ.js +4 -0
- dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
- dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
- dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
- dw/server/ui/assets/mdx-DAdMi_0p.js +1 -0
- dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
- dw/server/ui/assets/monaco--ixms01u.css +1 -0
- dw/server/ui/assets/monaco-BGCeEqaw.js +56 -0
- dw/server/ui/assets/msdax-DauUninz.js +1 -0
- dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
- dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
- dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
- dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
- dw/server/ui/assets/perl-oz_6vUea.js +1 -0
- dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
- dw/server/ui/assets/php-nr791fC2.js +1 -0
- dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
- dw/server/ui/assets/postiats-43DmfD33.js +1 -0
- dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
- dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
- dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
- dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
- dw/server/ui/assets/python-Bcn70HdC.js +1 -0
- dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
- dw/server/ui/assets/r-BwWrilGY.js +1 -0
- dw/server/ui/assets/razor-D1HmNnby.js +1 -0
- dw/server/ui/assets/redis-ClamHrr6.js +1 -0
- dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
- dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
- dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
- dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
- dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
- dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
- dw/server/ui/assets/scheme-BeGwcela.js +1 -0
- dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
- dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
- dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
- dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
- dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
- dw/server/ui/assets/sql-NEE52Syq.js +1 -0
- dw/server/ui/assets/st-DbInun42.js +1 -0
- dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
- dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
- dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
- dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
- dw/server/ui/assets/tsMode-D6u0XmOW.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-BU6v-LMV.js +1 -0
- dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
- dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
- dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
- dw/server/ui/assets/workers-Cn7cTUKr.js +1 -0
- dw/server/ui/assets/xml--0LP2Lwk.js +1 -0
- dw/server/ui/assets/yaml-mpBg9jnt.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/server/updater.py +192 -0
- dw/settings.py +98 -0
- dw/shot_span_preflight.py +116 -0
- dw/shots.py +359 -0
- dw/slice_preflight.py +148 -0
- dw/step.py +187 -0
- dw/step_cache.py +442 -0
- dw/subfolders.py +107 -0
- dw/task_domains.py +307 -0
- dw/tasks/assess.py +826 -0
- dw/tasks/audio_transcription.py +88 -0
- dw/tasks/audio_utils.py +1862 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/compose_text.py +74 -0
- dw/tasks/concat_videos.py +300 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/dissolve_videos.py +342 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +173 -0
- dw/tasks/grade.py +97 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +764 -0
- dw/tasks/interpolate_frames.py +252 -0
- dw/tasks/judge.py +68 -0
- dw/tasks/model_cache.py +55 -0
- dw/tasks/pair_audio.py +268 -0
- dw/tasks/qr_code.py +19 -0
- dw/tasks/restore_faces.py +175 -0
- dw/tasks/rife_model.py +192 -0
- dw/tasks/segment.py +121 -0
- dw/tasks/select.py +111 -0
- dw/tasks/speech_generation.py +228 -0
- dw/tasks/stabilize.py +129 -0
- dw/tasks/task.py +920 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +169 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +624 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +381 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +231 -0
- dw/validate.py +68 -0
- dw/variable_constraints.py +444 -0
- dw/variables.py +443 -0
- dw/video_extensions.py +141 -0
- dw/vram_estimate.py +116 -0
- dw/worker.py +764 -0
- dw/workflow.py +2007 -0
- dw/workflow_schema.json +1346 -0
- dw/workflow_sources.py +383 -0
- dw/workflows/h3_context_ir.json +57 -0
- dw/workflows/test.json +31 -0
- dw/workspace.py +730 -0
- dw_mcp/__init__.py +6 -0
- dw_mcp/__main__.py +133 -0
- dw_mcp/assets.py +336 -0
- dw_mcp/authoring.py +114 -0
- dw_mcp/catalog.py +360 -0
- dw_mcp/client.py +486 -0
- dw_mcp/diagnose.py +371 -0
- dw_mcp/exports.py +84 -0
- dw_mcp/guides.py +35 -0
- dw_mcp/media.py +638 -0
- dw_mcp/models.py +97 -0
- dw_mcp/prompts.py +104 -0
- dw_mcp/server.py +1343 -0
- 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
|