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
dw/variables.py
ADDED
|
@@ -0,0 +1,153 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import PIL
|
|
3
|
+
from .security import (
|
|
4
|
+
validate_variable_name,
|
|
5
|
+
validate_string_input,
|
|
6
|
+
SecurityError,
|
|
7
|
+
MAX_VARIABLE_VALUE_LENGTH,
|
|
8
|
+
)
|
|
9
|
+
|
|
10
|
+
logger = logging.getLogger("dw")
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def replace_variables(data, variables):
|
|
14
|
+
"""
|
|
15
|
+
Recursively replaces variable references in data structures with their actual values
|
|
16
|
+
Args:
|
|
17
|
+
data: The data structure (dict or list) containing variable references
|
|
18
|
+
variables: Dictionary of variable names and their values
|
|
19
|
+
"""
|
|
20
|
+
if variables is not None:
|
|
21
|
+
logger.debug(f"Processing variables: {list(variables.keys())}")
|
|
22
|
+
|
|
23
|
+
# Handle lists - replace any "variable:name" strings with their values
|
|
24
|
+
if isinstance(data, list):
|
|
25
|
+
logger.debug(f"Processing list of length {len(data)}")
|
|
26
|
+
for i, item in enumerate(data):
|
|
27
|
+
# Check for variable reference format "variable:name"
|
|
28
|
+
if isinstance(item, str) and item.startswith("variable:"):
|
|
29
|
+
variable_name = item.removeprefix("variable:")
|
|
30
|
+
logger.debug(f"Replacing variable reference: {variable_name}")
|
|
31
|
+
if not variable_name in variables:
|
|
32
|
+
logger.error(f"Variable <{variable_name}> not found")
|
|
33
|
+
raise Exception(f"Variable <{variable_name}> not found")
|
|
34
|
+
data[i] = variables[variable_name]
|
|
35
|
+
else:
|
|
36
|
+
# Recursively process nested structures
|
|
37
|
+
replace_variables(item, variables)
|
|
38
|
+
|
|
39
|
+
# Handle dictionaries - replace values that are variable references
|
|
40
|
+
elif isinstance(data, dict):
|
|
41
|
+
logger.debug(f"Processing dictionary with keys: {list(data.keys())}")
|
|
42
|
+
for k, v in data.items():
|
|
43
|
+
if isinstance(v, str) and v.startswith("variable:"):
|
|
44
|
+
variable_name = v.removeprefix("variable:")
|
|
45
|
+
logger.debug(f"Replacing variable reference: {variable_name}")
|
|
46
|
+
if not variable_name in variables:
|
|
47
|
+
logger.error(f"Variable <{variable_name}> not found")
|
|
48
|
+
raise Exception(f"Variable <{variable_name}> not found")
|
|
49
|
+
data[k] = variables[variable_name]
|
|
50
|
+
else:
|
|
51
|
+
# Recursively process nested structures in dictionary values
|
|
52
|
+
replace_variables(v, variables)
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def set_variables(values, variables):
|
|
56
|
+
"""
|
|
57
|
+
Sets the values of variables from a dictionary of new values with validation
|
|
58
|
+
Args:
|
|
59
|
+
values: Dictionary of new values to set
|
|
60
|
+
variables: Dictionary of existing variables with their default values/types
|
|
61
|
+
"""
|
|
62
|
+
logger.debug(f"Setting variables: {list(values.keys())}")
|
|
63
|
+
|
|
64
|
+
if not isinstance(values, dict) or not isinstance(variables, dict):
|
|
65
|
+
logger.error("Both values and variables must be dictionaries")
|
|
66
|
+
raise TypeError("Both values and variables must be dictionaries")
|
|
67
|
+
|
|
68
|
+
for k, v in values.items():
|
|
69
|
+
try:
|
|
70
|
+
# Validate variable name
|
|
71
|
+
validated_name = validate_variable_name(k)
|
|
72
|
+
|
|
73
|
+
# The workflow must have already declared this variable (with a default
|
|
74
|
+
# value/type) - reject unknown names instead of raising a bare KeyError
|
|
75
|
+
if validated_name not in variables:
|
|
76
|
+
declared = ", ".join(sorted(variables.keys()))
|
|
77
|
+
logger.error(
|
|
78
|
+
f"Unknown variable '{validated_name}'; declared variables: {declared}"
|
|
79
|
+
)
|
|
80
|
+
raise ValueError(
|
|
81
|
+
f"Unknown variable '{validated_name}'; declared variables: {declared}"
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
# Validate string values
|
|
85
|
+
if isinstance(v, str):
|
|
86
|
+
validated_value = validate_string_input(
|
|
87
|
+
v, max_length=MAX_VARIABLE_VALUE_LENGTH, allow_empty=True
|
|
88
|
+
)
|
|
89
|
+
else:
|
|
90
|
+
validated_value = v
|
|
91
|
+
|
|
92
|
+
logger.debug(
|
|
93
|
+
f"Setting variable {validated_name} to value: {validated_value}"
|
|
94
|
+
)
|
|
95
|
+
# Use the type of the existing variable to convert the new value
|
|
96
|
+
variables[validated_name] = get_value(
|
|
97
|
+
validated_value, type(variables[validated_name]), validated_name
|
|
98
|
+
)
|
|
99
|
+
|
|
100
|
+
except SecurityError as e:
|
|
101
|
+
logger.error(f"Security validation failed for variable {k}: {e}")
|
|
102
|
+
raise
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def get_value(v, desired_type, name=None):
|
|
106
|
+
"""
|
|
107
|
+
Converts a value to the desired type, with special handling for booleans
|
|
108
|
+
Args:
|
|
109
|
+
v: Value to convert
|
|
110
|
+
desired_type: Target type for conversion
|
|
111
|
+
name: Name of the variable being converted, used for error messages
|
|
112
|
+
Returns:
|
|
113
|
+
Converted value, or original value if conversion fails
|
|
114
|
+
"""
|
|
115
|
+
logger.debug(f"Converting value {v} to type {desired_type}")
|
|
116
|
+
|
|
117
|
+
# A variable declared null is an optional one the workflow states no type
|
|
118
|
+
# for - passing a value to it is the expected case, not a suspicious one
|
|
119
|
+
if desired_type is None or desired_type is type(None):
|
|
120
|
+
logger.debug("Variable has no declared type, using the value as given")
|
|
121
|
+
return v
|
|
122
|
+
|
|
123
|
+
# Special handling for boolean string values - bool("0") and bool("no") are
|
|
124
|
+
# both truthy in Python, which would silently invert the user's intent, so
|
|
125
|
+
# only a known set of true/false spellings is accepted here
|
|
126
|
+
if isinstance(v, str) and desired_type is bool:
|
|
127
|
+
lowered = v.lower()
|
|
128
|
+
if lowered in ("true", "1", "yes", "on"):
|
|
129
|
+
return True
|
|
130
|
+
if lowered in ("false", "0", "no", "off"):
|
|
131
|
+
return False
|
|
132
|
+
var_label = name if name is not None else "<unknown>"
|
|
133
|
+
message = f"Cannot interpret '{v}' as true/false for variable '{var_label}'"
|
|
134
|
+
logger.error(message)
|
|
135
|
+
raise ValueError(message)
|
|
136
|
+
|
|
137
|
+
# Special handling for list string values - list("cat") would mangle the
|
|
138
|
+
# string into ['c', 'a', 't'], so a comma-separated string is split instead
|
|
139
|
+
if isinstance(v, str) and desired_type is list:
|
|
140
|
+
return [item.strip() for item in v.split(",")]
|
|
141
|
+
|
|
142
|
+
# special handling for images that have already been realized
|
|
143
|
+
if isinstance(v, PIL.Image.Image):
|
|
144
|
+
return v
|
|
145
|
+
|
|
146
|
+
# Attempt type conversion, return original value if it fails
|
|
147
|
+
try:
|
|
148
|
+
converted = desired_type(v)
|
|
149
|
+
logger.debug(f"Successfully converted to {desired_type.__name__}: {converted}")
|
|
150
|
+
return converted
|
|
151
|
+
except Exception as e:
|
|
152
|
+
logger.warning(f"Failed to convert to {desired_type.__name__}: {e}")
|
|
153
|
+
return v
|
dw/worker.py
ADDED
|
@@ -0,0 +1,517 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Persistent worker process for workflow execution.
|
|
3
|
+
Keeps models loaded in GPU memory across multiple runs.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
import os
|
|
7
|
+
import sys
|
|
8
|
+
import queue
|
|
9
|
+
import logging
|
|
10
|
+
import threading
|
|
11
|
+
import traceback
|
|
12
|
+
from typing import Dict, Any
|
|
13
|
+
|
|
14
|
+
# Add parent directory to path for imports
|
|
15
|
+
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
|
16
|
+
|
|
17
|
+
from dw.workflow import workflow_from_file, workflow_from_definition
|
|
18
|
+
from dw.log_setup import setup_logging, set_log_level
|
|
19
|
+
from dw.settings import load_settings, resolve_path
|
|
20
|
+
from dw.security import validate_output_path
|
|
21
|
+
from dw.events import RunContext, WorkflowCancelled
|
|
22
|
+
from dw import get_device_type, empty_device_cache, device_memory_stats
|
|
23
|
+
|
|
24
|
+
logger = logging.getLogger("dw.worker")
|
|
25
|
+
|
|
26
|
+
# Memory management constants
|
|
27
|
+
MEMORY_GROWTH_THRESHOLD_MB = 500 # Warn if GPU memory grows by more than this
|
|
28
|
+
|
|
29
|
+
# How often (in seconds) the main loop wakes up to check whether the parent
|
|
30
|
+
# process is still alive when no command has arrived. Short enough that an
|
|
31
|
+
# orphaned worker exits promptly, long enough to avoid busy-waiting.
|
|
32
|
+
COMMAND_POLL_TIMEOUT_SECONDS = 5
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class WorkflowWorker:
|
|
36
|
+
"""
|
|
37
|
+
Persistent worker that keeps workflows and models loaded in memory.
|
|
38
|
+
Monitors workflow file for changes and reloads when necessary.
|
|
39
|
+
"""
|
|
40
|
+
|
|
41
|
+
def __init__(self, command_queue, result_queue, log_level="INFO"):
|
|
42
|
+
"""
|
|
43
|
+
Initialize the worker with communication queues.
|
|
44
|
+
|
|
45
|
+
Args:
|
|
46
|
+
command_queue: Queue for receiving commands from REPL
|
|
47
|
+
result_queue: Queue for sending results back to REPL
|
|
48
|
+
log_level: Logging level (DEBUG, INFO, WARNING, ERROR)
|
|
49
|
+
"""
|
|
50
|
+
self.command_queue = command_queue
|
|
51
|
+
self.result_queue = result_queue
|
|
52
|
+
|
|
53
|
+
# Capture the parent PID at startup so the main loop can detect
|
|
54
|
+
# orphaning (parent died/was killed without sending "shutdown") and
|
|
55
|
+
# exit cleanly instead of blocking forever on the command queue.
|
|
56
|
+
self.parent_pid = os.getppid()
|
|
57
|
+
|
|
58
|
+
# Log to the same file the rest of dw uses - the worker is a separate
|
|
59
|
+
# process, and ConcurrentRotatingFileHandler exists to share it safely
|
|
60
|
+
settings = load_settings()
|
|
61
|
+
setup_logging(resolve_path(settings.log_filename), log_level)
|
|
62
|
+
|
|
63
|
+
# Workflow state. Identity (path, or id for inline definitions) decides
|
|
64
|
+
# when the model cache is dropped wholesale - switching workflows frees
|
|
65
|
+
# the old one's models before the new one loads
|
|
66
|
+
self.workflow_identity = None
|
|
67
|
+
self.pending_shutdown = False
|
|
68
|
+
|
|
69
|
+
# Pipeline cache - persists across runs
|
|
70
|
+
self.loaded_pipelines = {}
|
|
71
|
+
self.shared_components = {}
|
|
72
|
+
|
|
73
|
+
# Memory tracking
|
|
74
|
+
self.run_count = 0
|
|
75
|
+
self.last_memory_mb = 0
|
|
76
|
+
|
|
77
|
+
logger.info("Worker process initialized")
|
|
78
|
+
|
|
79
|
+
def run(self):
|
|
80
|
+
"""
|
|
81
|
+
Main worker loop - processes commands until shutdown.
|
|
82
|
+
"""
|
|
83
|
+
logger.info("Worker entering command loop")
|
|
84
|
+
|
|
85
|
+
try:
|
|
86
|
+
while True:
|
|
87
|
+
try:
|
|
88
|
+
# Wait for a command from the REPL, but poll with a
|
|
89
|
+
# timeout rather than blocking forever. If nothing
|
|
90
|
+
# arrives, check whether the parent process is still
|
|
91
|
+
# alive - if it has died (e.g. crashed or was killed)
|
|
92
|
+
# without sending "shutdown", we'd otherwise sit here
|
|
93
|
+
# forever as an orphaned, unkillable-by-normal-means
|
|
94
|
+
# child. Exit cleanly instead.
|
|
95
|
+
try:
|
|
96
|
+
command = self.command_queue.get(
|
|
97
|
+
timeout=COMMAND_POLL_TIMEOUT_SECONDS
|
|
98
|
+
)
|
|
99
|
+
except queue.Empty:
|
|
100
|
+
if self._parent_is_dead():
|
|
101
|
+
logger.info(
|
|
102
|
+
f"Parent process (pid {self.parent_pid}) is gone - "
|
|
103
|
+
"worker exiting"
|
|
104
|
+
)
|
|
105
|
+
break
|
|
106
|
+
continue
|
|
107
|
+
|
|
108
|
+
command_type = command.get("type")
|
|
109
|
+
|
|
110
|
+
logger.debug(f"Received command: {command_type}")
|
|
111
|
+
|
|
112
|
+
if command_type == "execute":
|
|
113
|
+
self._handle_execute(command)
|
|
114
|
+
if self.pending_shutdown:
|
|
115
|
+
self._handle_shutdown()
|
|
116
|
+
break
|
|
117
|
+
elif command_type == "cancel":
|
|
118
|
+
# Nothing running - a cancel that raced the run's end
|
|
119
|
+
logger.debug("Ignoring cancel with no workflow running")
|
|
120
|
+
elif command_type == "shutdown":
|
|
121
|
+
self._handle_shutdown()
|
|
122
|
+
break
|
|
123
|
+
elif command_type == "ping":
|
|
124
|
+
self._handle_ping()
|
|
125
|
+
elif command_type == "clear_memory":
|
|
126
|
+
self._handle_clear_memory()
|
|
127
|
+
elif command_type == "memory_status":
|
|
128
|
+
self._handle_memory_status()
|
|
129
|
+
else:
|
|
130
|
+
self.result_queue.put(
|
|
131
|
+
{
|
|
132
|
+
"type": "error",
|
|
133
|
+
"message": f"Unknown command type: {command_type}",
|
|
134
|
+
}
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
except KeyboardInterrupt:
|
|
138
|
+
logger.info("Worker interrupted by keyboard")
|
|
139
|
+
break
|
|
140
|
+
except Exception as e:
|
|
141
|
+
logger.error(f"Error processing command: {e}", exc_info=True)
|
|
142
|
+
self.result_queue.put(
|
|
143
|
+
{
|
|
144
|
+
"type": "error",
|
|
145
|
+
"message": f"Command processing error: {str(e)}",
|
|
146
|
+
"traceback": traceback.format_exc(),
|
|
147
|
+
}
|
|
148
|
+
)
|
|
149
|
+
|
|
150
|
+
finally:
|
|
151
|
+
logger.info("Worker shutting down")
|
|
152
|
+
self._cleanup_all()
|
|
153
|
+
|
|
154
|
+
def _handle_execute(self, command: Dict[str, Any]):
|
|
155
|
+
"""
|
|
156
|
+
Execute a workflow, reusing loaded models if possible.
|
|
157
|
+
|
|
158
|
+
The command names the workflow either by path (workflow_path) or as an
|
|
159
|
+
inline definition (workflow, with an optional base_dir that relative
|
|
160
|
+
paths inside it resolve against). Models stay cached between runs of
|
|
161
|
+
the same workflow identity; pipelines are cached by what they load, so
|
|
162
|
+
an edited workflow keeps every pipeline whose definition is unchanged.
|
|
163
|
+
A {"type": "cancel"} command sent during execution stops the run at
|
|
164
|
+
the next step boundary or diffusion step.
|
|
165
|
+
|
|
166
|
+
Args:
|
|
167
|
+
command: Dictionary with workflow_path or workflow (+ base_dir),
|
|
168
|
+
arguments, output_dir, log_level
|
|
169
|
+
"""
|
|
170
|
+
arguments = command["arguments"]
|
|
171
|
+
output_dir = command["output_dir"]
|
|
172
|
+
log_level = command.get("log_level", "INFO")
|
|
173
|
+
|
|
174
|
+
try:
|
|
175
|
+
set_log_level(log_level)
|
|
176
|
+
|
|
177
|
+
workflow, identity = self._load_workflow(command, output_dir)
|
|
178
|
+
workflow.validate()
|
|
179
|
+
|
|
180
|
+
# Switching to a different workflow frees the old one's models
|
|
181
|
+
# before the new one loads - on one accelerator, holding both is
|
|
182
|
+
# what runs out of memory
|
|
183
|
+
if identity != self.workflow_identity:
|
|
184
|
+
if self.workflow_identity is not None:
|
|
185
|
+
self.result_queue.put(
|
|
186
|
+
{
|
|
187
|
+
"type": "output",
|
|
188
|
+
"message": "Workflow changed - releasing cached models...",
|
|
189
|
+
}
|
|
190
|
+
)
|
|
191
|
+
self._cleanup_all()
|
|
192
|
+
self.workflow_identity = identity
|
|
193
|
+
|
|
194
|
+
self.result_queue.put(
|
|
195
|
+
{"type": "workflow_loaded", "workflow_name": workflow.name}
|
|
196
|
+
)
|
|
197
|
+
self.result_queue.put(
|
|
198
|
+
{
|
|
199
|
+
"type": "output",
|
|
200
|
+
"message": f"Executing workflow: {workflow.name}",
|
|
201
|
+
}
|
|
202
|
+
)
|
|
203
|
+
|
|
204
|
+
# Progress events stream to the client as they happen; the watcher
|
|
205
|
+
# thread keeps the command queue live so cancel works mid-run
|
|
206
|
+
context = RunContext(
|
|
207
|
+
on_event=lambda event: self.result_queue.put(
|
|
208
|
+
{"type": "progress", **event}
|
|
209
|
+
)
|
|
210
|
+
)
|
|
211
|
+
watcher = self._watch_commands(context)
|
|
212
|
+
try:
|
|
213
|
+
workflow.run(arguments, self.loaded_pipelines, context=context)
|
|
214
|
+
finally:
|
|
215
|
+
watcher.stop()
|
|
216
|
+
|
|
217
|
+
# Drop cached pipelines this run no longer touched - an edited
|
|
218
|
+
# workflow that removed or redefined a step leaves those behind
|
|
219
|
+
for cache_key in list(self.loaded_pipelines):
|
|
220
|
+
if cache_key not in context.touched_pipelines:
|
|
221
|
+
logger.info("Evicting cached pipeline no longer in workflow")
|
|
222
|
+
del self.loaded_pipelines[cache_key]
|
|
223
|
+
|
|
224
|
+
self.run_count += 1
|
|
225
|
+
|
|
226
|
+
# Aggressive memory cleanup after execution
|
|
227
|
+
self._cleanup_between_runs()
|
|
228
|
+
|
|
229
|
+
# Report memory status
|
|
230
|
+
memory_info = self._get_memory_info()
|
|
231
|
+
self.result_queue.put({"type": "memory_info", "info": memory_info})
|
|
232
|
+
|
|
233
|
+
self.result_queue.put(
|
|
234
|
+
{
|
|
235
|
+
"type": "success",
|
|
236
|
+
"message": "Workflow completed successfully",
|
|
237
|
+
"run_count": self.run_count,
|
|
238
|
+
"manifest": getattr(workflow, "manifest", []),
|
|
239
|
+
}
|
|
240
|
+
)
|
|
241
|
+
|
|
242
|
+
except WorkflowCancelled:
|
|
243
|
+
self._cleanup_between_runs()
|
|
244
|
+
self.result_queue.put(
|
|
245
|
+
{"type": "cancelled", "message": "Workflow run cancelled"}
|
|
246
|
+
)
|
|
247
|
+
except Exception as e:
|
|
248
|
+
logger.error(f"Error executing workflow: {e}", exc_info=True)
|
|
249
|
+
self.result_queue.put(
|
|
250
|
+
{
|
|
251
|
+
"type": "error",
|
|
252
|
+
"message": f"Workflow execution error: {str(e)}",
|
|
253
|
+
"traceback": traceback.format_exc(),
|
|
254
|
+
}
|
|
255
|
+
)
|
|
256
|
+
|
|
257
|
+
def _load_workflow(self, command: Dict[str, Any], output_dir: str):
|
|
258
|
+
"""Build the Workflow a command names, and its cache identity."""
|
|
259
|
+
if "workflow_path" in command and command["workflow_path"] is not None:
|
|
260
|
+
workflow_path = command["workflow_path"]
|
|
261
|
+
workflow = workflow_from_file(workflow_path, output_dir)
|
|
262
|
+
return workflow, ("path", workflow_path)
|
|
263
|
+
|
|
264
|
+
workflow_data = command["workflow"]
|
|
265
|
+
workflow = workflow_from_definition(
|
|
266
|
+
workflow_data, output_dir, command.get("base_dir")
|
|
267
|
+
)
|
|
268
|
+
return workflow, ("inline", workflow_data.get("id"))
|
|
269
|
+
|
|
270
|
+
def _watch_commands(self, context):
|
|
271
|
+
"""Watch the command queue during a run so cancel and ping still work.
|
|
272
|
+
|
|
273
|
+
Returns an object with stop(); anything that is not cancel, ping or
|
|
274
|
+
shutdown is refused, since one workflow runs at a time.
|
|
275
|
+
"""
|
|
276
|
+
stop_event = threading.Event()
|
|
277
|
+
worker = self
|
|
278
|
+
|
|
279
|
+
def watch():
|
|
280
|
+
while not stop_event.is_set():
|
|
281
|
+
try:
|
|
282
|
+
command = worker.command_queue.get(timeout=0.25)
|
|
283
|
+
except queue.Empty:
|
|
284
|
+
continue
|
|
285
|
+
command_type = command.get("type")
|
|
286
|
+
if command_type == "cancel":
|
|
287
|
+
logger.info("Cancel requested")
|
|
288
|
+
context.cancel()
|
|
289
|
+
worker.result_queue.put(
|
|
290
|
+
{"type": "output", "message": "Cancelling..."}
|
|
291
|
+
)
|
|
292
|
+
elif command_type == "ping":
|
|
293
|
+
worker._handle_ping()
|
|
294
|
+
elif command_type == "shutdown":
|
|
295
|
+
# Stop the run, then let the main loop see the shutdown
|
|
296
|
+
context.cancel()
|
|
297
|
+
worker.pending_shutdown = True
|
|
298
|
+
else:
|
|
299
|
+
worker.result_queue.put(
|
|
300
|
+
{
|
|
301
|
+
"type": "error",
|
|
302
|
+
"message": f"Cannot handle '{command_type}' while a "
|
|
303
|
+
"workflow is running",
|
|
304
|
+
}
|
|
305
|
+
)
|
|
306
|
+
|
|
307
|
+
thread = threading.Thread(target=watch, daemon=True, name="command-watcher")
|
|
308
|
+
thread.start()
|
|
309
|
+
|
|
310
|
+
class _Watcher:
|
|
311
|
+
def stop(self):
|
|
312
|
+
stop_event.set()
|
|
313
|
+
thread.join()
|
|
314
|
+
|
|
315
|
+
return _Watcher()
|
|
316
|
+
|
|
317
|
+
def _handle_shutdown(self):
|
|
318
|
+
"""Handle graceful shutdown request."""
|
|
319
|
+
logger.info("Shutdown requested")
|
|
320
|
+
self._cleanup_all()
|
|
321
|
+
self.result_queue.put({"type": "shutdown_complete"})
|
|
322
|
+
|
|
323
|
+
def _handle_ping(self):
|
|
324
|
+
"""Respond to ping to prove worker is alive."""
|
|
325
|
+
self.result_queue.put({"type": "pong", "run_count": self.run_count})
|
|
326
|
+
|
|
327
|
+
def _handle_clear_memory(self):
|
|
328
|
+
"""Handle explicit memory clear request."""
|
|
329
|
+
logger.info("Memory clear requested")
|
|
330
|
+
self._cleanup_all()
|
|
331
|
+
memory_info = self._get_memory_info()
|
|
332
|
+
self.result_queue.put({"type": "memory_cleared", "info": memory_info})
|
|
333
|
+
|
|
334
|
+
def _handle_memory_status(self):
|
|
335
|
+
"""Report current memory usage."""
|
|
336
|
+
memory_info = self._get_memory_info()
|
|
337
|
+
self.result_queue.put({"type": "memory_status", "info": memory_info})
|
|
338
|
+
|
|
339
|
+
def _cleanup_between_runs(self):
|
|
340
|
+
"""
|
|
341
|
+
Aggressive memory cleanup between workflow runs.
|
|
342
|
+
Keeps models loaded but cleans up intermediate tensors and garbage.
|
|
343
|
+
"""
|
|
344
|
+
import gc
|
|
345
|
+
|
|
346
|
+
logger.debug("Performing inter-run memory cleanup")
|
|
347
|
+
|
|
348
|
+
# Force garbage collection
|
|
349
|
+
gc.collect()
|
|
350
|
+
|
|
351
|
+
# Clean up GPU cache if available (CUDA or MPS). Don't synchronize
|
|
352
|
+
# here as it's expensive and unnecessary.
|
|
353
|
+
try:
|
|
354
|
+
empty_device_cache()
|
|
355
|
+
except Exception as e:
|
|
356
|
+
logger.warning(f"Could not clean GPU cache: {e}")
|
|
357
|
+
|
|
358
|
+
# Check for memory growth
|
|
359
|
+
current_memory = self._get_gpu_memory_mb()
|
|
360
|
+
if current_memory > 0:
|
|
361
|
+
if self.last_memory_mb > 0:
|
|
362
|
+
growth = current_memory - self.last_memory_mb
|
|
363
|
+
if growth > MEMORY_GROWTH_THRESHOLD_MB:
|
|
364
|
+
logger.warning(
|
|
365
|
+
f"GPU memory grew by {growth:.1f}MB "
|
|
366
|
+
f"({self.last_memory_mb:.1f}MB -> {current_memory:.1f}MB)"
|
|
367
|
+
)
|
|
368
|
+
self.last_memory_mb = current_memory
|
|
369
|
+
|
|
370
|
+
logger.debug("Inter-run cleanup complete")
|
|
371
|
+
|
|
372
|
+
def _cleanup_all(self):
|
|
373
|
+
"""
|
|
374
|
+
Complete cleanup - clear all cached models and components.
|
|
375
|
+
Called when workflow changes or on shutdown.
|
|
376
|
+
"""
|
|
377
|
+
import gc
|
|
378
|
+
from .tasks.model_cache import clear_model_cache
|
|
379
|
+
|
|
380
|
+
logger.info("Performing full cleanup")
|
|
381
|
+
|
|
382
|
+
# Clear pipeline cache and any models task handlers cached
|
|
383
|
+
self.loaded_pipelines.clear()
|
|
384
|
+
self.shared_components.clear()
|
|
385
|
+
clear_model_cache()
|
|
386
|
+
|
|
387
|
+
# Reset state
|
|
388
|
+
self.run_count = 0
|
|
389
|
+
self.last_memory_mb = 0
|
|
390
|
+
|
|
391
|
+
# Force garbage collection multiple times
|
|
392
|
+
for _ in range(3):
|
|
393
|
+
gc.collect()
|
|
394
|
+
|
|
395
|
+
# Aggressive GPU cleanup (CUDA or MPS) - empty cache and synchronize
|
|
396
|
+
# to ensure all operations complete before we go on to reset stats.
|
|
397
|
+
try:
|
|
398
|
+
empty_device_cache(synchronize=True)
|
|
399
|
+
|
|
400
|
+
# Try to reset CUDA memory stats - no MPS equivalent exists
|
|
401
|
+
if get_device_type() == "cuda":
|
|
402
|
+
import torch
|
|
403
|
+
|
|
404
|
+
try:
|
|
405
|
+
torch.cuda.reset_peak_memory_stats()
|
|
406
|
+
torch.cuda.reset_accumulated_memory_stats()
|
|
407
|
+
except (RuntimeError, AttributeError) as e:
|
|
408
|
+
logger.debug(f"Could not reset memory stats: {e}")
|
|
409
|
+
|
|
410
|
+
except Exception as e:
|
|
411
|
+
logger.warning(f"Could not perform GPU cleanup: {e}")
|
|
412
|
+
|
|
413
|
+
logger.info("Full cleanup complete")
|
|
414
|
+
|
|
415
|
+
def _parent_is_dead(self) -> bool:
|
|
416
|
+
"""
|
|
417
|
+
Check whether the process that spawned this worker is still around.
|
|
418
|
+
|
|
419
|
+
On POSIX, a process gets reparented to init (traditionally pid 1,
|
|
420
|
+
though some systems use a subreaper) once its original parent exits,
|
|
421
|
+
so a changed getppid() is the standard signal that we've been
|
|
422
|
+
orphaned.
|
|
423
|
+
|
|
424
|
+
Returns:
|
|
425
|
+
True if the parent appears to be gone, False otherwise.
|
|
426
|
+
"""
|
|
427
|
+
current_ppid = os.getppid()
|
|
428
|
+
return current_ppid != self.parent_pid or current_ppid == 1
|
|
429
|
+
|
|
430
|
+
def _get_gpu_memory_mb(self) -> float:
|
|
431
|
+
"""
|
|
432
|
+
Get current GPU memory usage in MB.
|
|
433
|
+
|
|
434
|
+
Returns:
|
|
435
|
+
Memory usage in MB, or 0 if not available
|
|
436
|
+
"""
|
|
437
|
+
try:
|
|
438
|
+
return device_memory_stats()["allocated_mb"]
|
|
439
|
+
except (RuntimeError, AttributeError) as e:
|
|
440
|
+
logger.debug(f"Could not get GPU memory: {e}")
|
|
441
|
+
return 0.0
|
|
442
|
+
|
|
443
|
+
def _get_memory_info(self) -> Dict[str, Any]:
|
|
444
|
+
"""
|
|
445
|
+
Get detailed memory information.
|
|
446
|
+
|
|
447
|
+
Returns:
|
|
448
|
+
Dictionary with memory statistics
|
|
449
|
+
"""
|
|
450
|
+
info = {
|
|
451
|
+
"run_count": self.run_count,
|
|
452
|
+
"gpu_available": False,
|
|
453
|
+
"gpu_memory_allocated_mb": 0.0,
|
|
454
|
+
"gpu_memory_reserved_mb": 0.0,
|
|
455
|
+
"gpu_memory_free_mb": 0.0,
|
|
456
|
+
"gpu_device_name": None,
|
|
457
|
+
}
|
|
458
|
+
|
|
459
|
+
try:
|
|
460
|
+
stats = device_memory_stats()
|
|
461
|
+
info["gpu_available"] = stats["available"]
|
|
462
|
+
info["gpu_device_name"] = stats["device_name"]
|
|
463
|
+
info["gpu_memory_allocated_mb"] = stats["allocated_mb"]
|
|
464
|
+
info["gpu_memory_reserved_mb"] = stats["reserved_mb"]
|
|
465
|
+
# free/total are only ever None when CUDA's mem_get_info call
|
|
466
|
+
# itself failed - leave gpu_memory_free_mb at its 0.0 default and
|
|
467
|
+
# gpu_memory_total_mb unset in that case, same as before.
|
|
468
|
+
if stats["free_mb"] is not None:
|
|
469
|
+
info["gpu_memory_free_mb"] = stats["free_mb"]
|
|
470
|
+
if stats["total_mb"] is not None:
|
|
471
|
+
info["gpu_memory_total_mb"] = stats["total_mb"]
|
|
472
|
+
except (ImportError, RuntimeError, AttributeError) as e:
|
|
473
|
+
logger.debug(f"Could not access GPU: {e}")
|
|
474
|
+
|
|
475
|
+
return info
|
|
476
|
+
|
|
477
|
+
|
|
478
|
+
def worker_main(command_queue, result_queue, log_level="INFO"):
|
|
479
|
+
"""
|
|
480
|
+
Entry point for worker process.
|
|
481
|
+
|
|
482
|
+
Args:
|
|
483
|
+
command_queue: Queue for receiving commands
|
|
484
|
+
result_queue: Queue for sending results
|
|
485
|
+
log_level: Logging level
|
|
486
|
+
"""
|
|
487
|
+
try:
|
|
488
|
+
worker = WorkflowWorker(command_queue, result_queue, log_level)
|
|
489
|
+
worker.run()
|
|
490
|
+
except Exception as e:
|
|
491
|
+
logger.error(f"Worker crashed: {e}", exc_info=True)
|
|
492
|
+
try:
|
|
493
|
+
result_queue.put(
|
|
494
|
+
{
|
|
495
|
+
"type": "worker_crashed",
|
|
496
|
+
"message": str(e),
|
|
497
|
+
"traceback": traceback.format_exc(),
|
|
498
|
+
}
|
|
499
|
+
)
|
|
500
|
+
except (OSError, RuntimeError) as queue_error:
|
|
501
|
+
logger.error(f"Failed to send crash notification to queue: {queue_error}")
|
|
502
|
+
sys.exit(1)
|
|
503
|
+
|
|
504
|
+
|
|
505
|
+
if __name__ == "__main__":
|
|
506
|
+
# For testing - won't normally be run directly
|
|
507
|
+
import multiprocessing
|
|
508
|
+
|
|
509
|
+
# Set spawn method for CUDA compatibility
|
|
510
|
+
if multiprocessing.get_start_method(allow_none=True) != "spawn":
|
|
511
|
+
multiprocessing.set_start_method("spawn", force=True)
|
|
512
|
+
|
|
513
|
+
cmd_queue = multiprocessing.Queue()
|
|
514
|
+
res_queue = multiprocessing.Queue()
|
|
515
|
+
|
|
516
|
+
print("Starting worker in test mode...")
|
|
517
|
+
worker_main(cmd_queue, res_queue, "DEBUG")
|