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.
- diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
- diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
- diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
- diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
- dw/__init__.py +353 -0
- dw/arguments.py +906 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +145 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/events.py +78 -0
- dw/hub_cache.py +289 -0
- dw/introspection.py +458 -0
- dw/log_setup.py +45 -0
- dw/pipeline_processors/chain.py +750 -0
- dw/pipeline_processors/config_objects.py +235 -0
- dw/pipeline_processors/pipeline.py +1687 -0
- dw/pipeline_processors/remote.py +18 -0
- dw/previous_results.py +259 -0
- dw/prompt_weighting.py +378 -0
- dw/repl.py +298 -0
- dw/repl_commands.py +808 -0
- dw/repl_worker.py +129 -0
- dw/result.py +850 -0
- dw/run.py +92 -0
- dw/schema.py +24 -0
- dw/security.py +379 -0
- dw/serve.py +70 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +588 -0
- dw/server/jobs.py +547 -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-CEh6hWi2.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-CExg3_mM.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-DH6orYh2.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-CbrMVW4Q.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-YDNPZw2M.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
- dw/server/ui/assets/index-B7-VcYS-.css +1 -0
- dw/server/ui/assets/index-D_EiPU3b.js +13 -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-dYuBvioq.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-CUqLM39V.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-D6vxBzMv.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.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-SHQb6vmD.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-CP-s5rcP.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-x0_EGHq9.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-BZC4LQDP.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-BTfA6SbD.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-CWA4MsNk.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-CWU0uvj5.js +1 -0
- dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
- dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/settings.py +77 -0
- dw/step.py +132 -0
- dw/tasks/audio_utils.py +266 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/concat_videos.py +80 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +139 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +661 -0
- dw/tasks/interpolate_frames.py +227 -0
- dw/tasks/model_cache.py +39 -0
- dw/tasks/pair_audio.py +58 -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/task.py +474 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +168 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +154 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +376 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +68 -0
- dw/validate.py +43 -0
- dw/variables.py +153 -0
- dw/worker.py +517 -0
- dw/workflow.py +553 -0
- dw/workflow_schema.json +1157 -0
- dw/workflows/augment_prompt.json +65 -0
- dw/workflows/describe_image.json +58 -0
- dw/workflows/h3_context_ir.json +57 -0
- 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
|