diffusers-workflow 0.4.0a3__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (171) hide show
  1. diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
  2. diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
  3. diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
  4. diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
  5. diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
  6. diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
  7. dw/__init__.py +353 -0
  8. dw/arguments.py +906 -0
  9. dw/cache_blocks.json +16 -0
  10. dw/cache_blocks.py +145 -0
  11. dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
  12. dw/events.py +78 -0
  13. dw/hub_cache.py +289 -0
  14. dw/introspection.py +458 -0
  15. dw/log_setup.py +45 -0
  16. dw/pipeline_processors/chain.py +750 -0
  17. dw/pipeline_processors/config_objects.py +235 -0
  18. dw/pipeline_processors/pipeline.py +1687 -0
  19. dw/pipeline_processors/remote.py +18 -0
  20. dw/previous_results.py +259 -0
  21. dw/prompt_weighting.py +378 -0
  22. dw/repl.py +298 -0
  23. dw/repl_commands.py +808 -0
  24. dw/repl_worker.py +129 -0
  25. dw/result.py +850 -0
  26. dw/run.py +92 -0
  27. dw/schema.py +24 -0
  28. dw/security.py +379 -0
  29. dw/serve.py +70 -0
  30. dw/server/__init__.py +2 -0
  31. dw/server/app.py +588 -0
  32. dw/server/jobs.py +547 -0
  33. dw/server/ui/assets/abap-08VXUWAP.js +1 -0
  34. dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
  35. dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
  36. dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
  37. dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
  38. dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
  39. dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
  40. dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
  41. dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
  42. dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
  43. dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
  44. dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
  45. dw/server/ui/assets/css-DIMkf-bt.js +3 -0
  46. dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
  47. dw/server/ui/assets/cssMode-CEh6hWi2.js +1 -0
  48. dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
  49. dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
  50. dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
  51. dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
  52. dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
  53. dw/server/ui/assets/editor.api-CExg3_mM.js +847 -0
  54. dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
  55. dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
  56. dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
  57. dw/server/ui/assets/freemarker2-DH6orYh2.js +3 -0
  58. dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
  59. dw/server/ui/assets/go-C-y9NEjX.js +1 -0
  60. dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
  61. dw/server/ui/assets/handlebars-CbrMVW4Q.js +1 -0
  62. dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
  63. dw/server/ui/assets/html-YDNPZw2M.js +1 -0
  64. dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
  65. dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
  66. dw/server/ui/assets/index-B7-VcYS-.css +1 -0
  67. dw/server/ui/assets/index-D_EiPU3b.js +13 -0
  68. dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
  69. dw/server/ui/assets/java-BEtHBSE6.js +1 -0
  70. dw/server/ui/assets/javascript-dYuBvioq.js +1 -0
  71. dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
  72. dw/server/ui/assets/jsonMode-CUqLM39V.js +7 -0
  73. dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
  74. dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
  75. dw/server/ui/assets/less-B9JPFI3C.js +2 -0
  76. dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
  77. dw/server/ui/assets/liquid-D6vxBzMv.js +1 -0
  78. dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.js +4 -0
  79. dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
  80. dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
  81. dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
  82. dw/server/ui/assets/mdx-SHQb6vmD.js +1 -0
  83. dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
  84. dw/server/ui/assets/monaco--ixms01u.css +1 -0
  85. dw/server/ui/assets/monaco-CP-s5rcP.js +56 -0
  86. dw/server/ui/assets/msdax-DauUninz.js +1 -0
  87. dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
  88. dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
  89. dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
  90. dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
  91. dw/server/ui/assets/perl-oz_6vUea.js +1 -0
  92. dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
  93. dw/server/ui/assets/php-nr791fC2.js +1 -0
  94. dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
  95. dw/server/ui/assets/postiats-43DmfD33.js +1 -0
  96. dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
  97. dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
  98. dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
  99. dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
  100. dw/server/ui/assets/python-x0_EGHq9.js +1 -0
  101. dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
  102. dw/server/ui/assets/r-BwWrilGY.js +1 -0
  103. dw/server/ui/assets/razor-BZC4LQDP.js +1 -0
  104. dw/server/ui/assets/redis-ClamHrr6.js +1 -0
  105. dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
  106. dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
  107. dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
  108. dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
  109. dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
  110. dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
  111. dw/server/ui/assets/scheme-BeGwcela.js +1 -0
  112. dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
  113. dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
  114. dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
  115. dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
  116. dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
  117. dw/server/ui/assets/sql-NEE52Syq.js +1 -0
  118. dw/server/ui/assets/st-DbInun42.js +1 -0
  119. dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
  120. dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
  121. dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
  122. dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
  123. dw/server/ui/assets/tsMode-BTfA6SbD.js +11 -0
  124. dw/server/ui/assets/twig-De2hgUGE.js +1 -0
  125. dw/server/ui/assets/typescript-CWA4MsNk.js +1 -0
  126. dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
  127. dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
  128. dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
  129. dw/server/ui/assets/workers-CWU0uvj5.js +1 -0
  130. dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
  131. dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
  132. dw/server/ui/index.html +17 -0
  133. dw/settings.py +77 -0
  134. dw/step.py +132 -0
  135. dw/tasks/audio_utils.py +266 -0
  136. dw/tasks/background_remover.py +43 -0
  137. dw/tasks/borders.py +113 -0
  138. dw/tasks/concat_videos.py +80 -0
  139. dw/tasks/depth_estimator.py +54 -0
  140. dw/tasks/diffusion_upscale.py +109 -0
  141. dw/tasks/format_messages.py +24 -0
  142. dw/tasks/gather.py +139 -0
  143. dw/tasks/image_to_text.py +43 -0
  144. dw/tasks/image_utils.py +661 -0
  145. dw/tasks/interpolate_frames.py +227 -0
  146. dw/tasks/model_cache.py +39 -0
  147. dw/tasks/pair_audio.py +58 -0
  148. dw/tasks/qr_code.py +19 -0
  149. dw/tasks/restore_faces.py +175 -0
  150. dw/tasks/rife_model.py +192 -0
  151. dw/tasks/segment.py +121 -0
  152. dw/tasks/task.py +474 -0
  153. dw/tasks/tensor_image.py +57 -0
  154. dw/tasks/text_generation.py +168 -0
  155. dw/tasks/text_sections.py +80 -0
  156. dw/tasks/upscale.py +203 -0
  157. dw/tasks/video_utils.py +154 -0
  158. dw/tasks/zoe_depth.py +71 -0
  159. dw/teacache.py +376 -0
  160. dw/teacache_models.json +99 -0
  161. dw/test.py +29 -0
  162. dw/type_helpers.py +68 -0
  163. dw/validate.py +43 -0
  164. dw/variables.py +153 -0
  165. dw/worker.py +517 -0
  166. dw/workflow.py +553 -0
  167. dw/workflow_schema.json +1157 -0
  168. dw/workflows/augment_prompt.json +65 -0
  169. dw/workflows/describe_image.json +58 -0
  170. dw/workflows/h3_context_ir.json +57 -0
  171. dw/workflows/test.json +31 -0
dw/tasks/task.py ADDED
@@ -0,0 +1,474 @@
1
+ import logging
2
+ from typing import Callable, Dict
3
+ from .qr_code import get_qrcode_image
4
+ from .image_utils import process_image
5
+ from .video_utils import process_video
6
+ from .gather import gather_images, gather_inputs, gather_videos
7
+ from .format_messages import (
8
+ format_chat_message,
9
+ batch_decode_post_process,
10
+ get_dict_value,
11
+ )
12
+
13
+ # The model-backed handlers (upscale, restore_faces, segment, interpolate_frames,
14
+ # image_to_text, text_generation, diffusion_upscale) are imported inside their
15
+ # handlers - at module scope their transformers/model imports add seconds to
16
+ # every startup for workflows that never run those tasks
17
+
18
+ logger = logging.getLogger("dw")
19
+
20
+
21
+ # Command registry: maps command names to handler functions
22
+ _COMMAND_REGISTRY: Dict[str, Callable] = {}
23
+
24
+ # What each command's arguments actually are. The handlers forward
25
+ # **arguments into an implementation function, so that function's signature
26
+ # is the command's argument schema - registering its dotted path here (a
27
+ # string, to preserve the lazy-import discipline) lets the introspection
28
+ # layer read the same signature the runtime calls. 'provided' names the
29
+ # parameters the dispatch supplies itself, which are not workflow arguments.
30
+ _COMMAND_INFO: Dict[str, dict] = {}
31
+
32
+
33
+ def register_command(command_name: str, implementation=None, provided=()):
34
+ """
35
+ Decorator to register a command handler function.
36
+
37
+ Args:
38
+ command_name: The command name to register
39
+ implementation: Dotted path to the function whose signature defines
40
+ the command's arguments (None for a command that consumes a
41
+ free-form dict)
42
+ provided: Parameter names the dispatch supplies itself
43
+
44
+ Returns:
45
+ Decorator function
46
+ """
47
+
48
+ def decorator(func: Callable) -> Callable:
49
+ _COMMAND_REGISTRY[command_name] = func
50
+ _COMMAND_INFO[command_name] = {
51
+ "kind": "command",
52
+ "implementation": implementation,
53
+ "provided": tuple(provided),
54
+ }
55
+ logger.debug(f"Registered command handler: {command_name}")
56
+ return func
57
+
58
+ return decorator
59
+
60
+
61
+ def task_command_info(command_name):
62
+ """Where a task command's argument schema lives: a dict with 'kind'
63
+ ('command', 'image_processor' or 'video_processor'), 'implementation'
64
+ (dotted path or None for free-form), and 'provided'. Raises ValueError
65
+ for a name that is not a task command at all."""
66
+ info = _COMMAND_INFO.get(command_name)
67
+ if info is not None:
68
+ return info
69
+ if command_name in _VIDEO_PROCESSOR_INFO:
70
+ return _VIDEO_PROCESSOR_INFO[command_name]
71
+ from .image_utils import available_processors
72
+
73
+ if command_name in available_processors():
74
+ return {"kind": "image_processor", "implementation": None, "provided": ()}
75
+ raise ValueError(f"Unknown task command: '{command_name}'")
76
+
77
+
78
+ # Command handler functions
79
+ @register_command("qr_code", implementation="dw.tasks.qr_code.get_qrcode_image")
80
+ def _handle_qr_code(task, arguments, previous_pipelines):
81
+ """Generate QR code image"""
82
+ logger.debug("Generating QR code")
83
+ return get_qrcode_image(**arguments)
84
+
85
+
86
+ @register_command("gather_images", implementation="dw.tasks.gather.gather_images")
87
+ def _handle_gather_images(task, arguments, previous_pipelines):
88
+ """Gather multiple images"""
89
+ logger.debug("Gathering images")
90
+ return gather_images(**arguments)
91
+
92
+
93
+ @register_command("gather_videos", implementation="dw.tasks.gather.gather_videos")
94
+ def _handle_gather_videos(task, arguments, previous_pipelines):
95
+ """Gather multiple videos"""
96
+ logger.debug("Gathering videos")
97
+ return gather_videos(**arguments)
98
+
99
+
100
+ # gather_inputs passes its whole dict through unchanged - free-form by design
101
+ @register_command("gather_inputs")
102
+ def _handle_gather_inputs(task, arguments, previous_pipelines):
103
+ """Gather inputs from various sources"""
104
+ logger.debug("Gathering inputs")
105
+ return gather_inputs(arguments)
106
+
107
+
108
+ @register_command(
109
+ "concat_videos", implementation="dw.tasks.concat_videos.concat_videos"
110
+ )
111
+ def _handle_concat_videos(task, arguments, previous_pipelines):
112
+ """Concatenate videos - and the audio generated with them - into one"""
113
+ logger.debug("Concatenating videos")
114
+ from .concat_videos import concat_videos
115
+
116
+ return concat_videos(**arguments)
117
+
118
+
119
+ @register_command("slice_audio", implementation="dw.tasks.audio_utils.slice_audio")
120
+ def _handle_slice_audio(task, arguments, previous_pipelines):
121
+ """Cut a time- or frame-aligned slice out of an audio track"""
122
+ logger.debug("Slicing audio")
123
+ from .audio_utils import slice_audio
124
+
125
+ return slice_audio(**arguments)
126
+
127
+
128
+ @register_command("video_frames", implementation="dw.tasks.video_utils.frames_as_array")
129
+ def _handle_video_frames(task, arguments, previous_pipelines):
130
+ """The frames of a generated video, as one array a later step can condition on"""
131
+ logger.debug("Extracting video frames")
132
+ from .video_utils import frames_as_array
133
+
134
+ return frames_as_array(**arguments)
135
+
136
+
137
+ @register_command("pair_audio", implementation="dw.tasks.pair_audio.pair_audio")
138
+ def _handle_pair_audio(task, arguments, previous_pipelines):
139
+ """Pair a video's frames with an audio track generated beside them"""
140
+ logger.debug("Pairing audio with video")
141
+ from .pair_audio import pair_audio
142
+
143
+ return pair_audio(**arguments)
144
+
145
+
146
+ @register_command(
147
+ "crossfade_audio", implementation="dw.tasks.audio_utils.crossfade_audio"
148
+ )
149
+ def _handle_crossfade_audio(task, arguments, previous_pipelines):
150
+ """Join audio tracks with an equal-power crossfade"""
151
+ logger.debug("Crossfading audio")
152
+ from .audio_utils import crossfade_audio
153
+
154
+ return crossfade_audio(**arguments)
155
+
156
+
157
+ @register_command(
158
+ "format_chat_message", implementation="dw.tasks.format_messages.format_chat_message"
159
+ )
160
+ def _handle_format_chat_message(task, arguments, previous_pipelines):
161
+ """Format chat message for LLM input"""
162
+ logger.debug("Formatting chat message")
163
+ return format_chat_message(**arguments)
164
+
165
+
166
+ @register_command(
167
+ "get_dict_value", implementation="dw.tasks.format_messages.get_dict_value"
168
+ )
169
+ def _handle_get_dict_value(task, arguments, previous_pipelines):
170
+ """Extract value from dictionary"""
171
+ logger.debug("Getting dictionary value")
172
+ return get_dict_value(**arguments)
173
+
174
+
175
+ @register_command("upscale", implementation="dw.tasks.upscale.upscale_image")
176
+ def _handle_upscale(task, arguments, previous_pipelines):
177
+ """Upscale an image using a spandrel-compatible super-resolution model"""
178
+ logger.debug("Upscaling image")
179
+ image = arguments.pop("image")
180
+ model_name = arguments.pop("model_name")
181
+ from .upscale import upscale_image
182
+
183
+ return upscale_image(
184
+ image, model_name, device=task.device_for(arguments), **arguments
185
+ )
186
+
187
+
188
+ @register_command(
189
+ "diffusion_upscale", implementation="dw.tasks.diffusion_upscale.diffusion_upscale"
190
+ )
191
+ def _handle_diffusion_upscale(task, arguments, previous_pipelines):
192
+ """Upscale an image using a diffusion-based upscale pipeline"""
193
+ logger.debug("Diffusion upscaling image")
194
+ image = arguments.pop("image")
195
+ from .diffusion_upscale import diffusion_upscale
196
+
197
+ return diffusion_upscale(image, device=task.device_for(arguments), **arguments)
198
+
199
+
200
+ @register_command(
201
+ "restore_faces", implementation="dw.tasks.restore_faces.restore_faces"
202
+ )
203
+ def _handle_restore_faces(task, arguments, previous_pipelines):
204
+ """Restore faces in an image using a spandrel-compatible face restoration model"""
205
+ logger.debug("Restoring faces")
206
+ image = arguments.pop("image")
207
+ model_name = arguments.pop("model_name")
208
+ from .restore_faces import restore_faces
209
+
210
+ return restore_faces(
211
+ image, model_name, device=task.device_for(arguments), **arguments
212
+ )
213
+
214
+
215
+ @register_command("segment", implementation="dw.tasks.segment.segment_image")
216
+ def _handle_segment(task, arguments, previous_pipelines):
217
+ """Segment objects in an image using text prompt"""
218
+ logger.debug("Segmenting image")
219
+ image = arguments.pop("image")
220
+ prompt = arguments.pop("prompt")
221
+ from .segment import segment_image
222
+
223
+ return segment_image(image, prompt, device=task.device_for(arguments), **arguments)
224
+
225
+
226
+ @register_command(
227
+ "interpolate_frames",
228
+ implementation="dw.tasks.interpolate_frames.interpolate_frames",
229
+ )
230
+ def _handle_interpolate_frames(task, arguments, previous_pipelines):
231
+ """Interpolate video frames to increase frame rate"""
232
+ logger.debug("Interpolating frames")
233
+ video = arguments.pop("video")
234
+ from .interpolate_frames import interpolate_frames
235
+
236
+ return interpolate_frames(video, device=task.device_for(arguments), **arguments)
237
+
238
+
239
+ @register_command(
240
+ "image_to_text", implementation="dw.tasks.image_to_text.image_to_text"
241
+ )
242
+ def _handle_image_to_text(task, arguments, previous_pipelines):
243
+ """Generate text caption from an image"""
244
+ logger.debug("Captioning image")
245
+ image = arguments.pop("image")
246
+ from .image_to_text import image_to_text
247
+
248
+ return image_to_text(image, device=task.device_for(arguments), **arguments)
249
+
250
+
251
+ @register_command(
252
+ "text_generation", implementation="dw.tasks.text_generation.generate_text"
253
+ )
254
+ def _handle_text_generation(task, arguments, previous_pipelines):
255
+ """Generate text from a prompt using a local LLM"""
256
+ logger.debug("Generating text")
257
+ prompt = arguments.pop("prompt")
258
+ from .text_generation import generate_text
259
+
260
+ return generate_text(prompt, device=task.device_for(arguments), **arguments)
261
+
262
+
263
+ @register_command(
264
+ "extract_sections", implementation="dw.tasks.text_sections.extract_sections"
265
+ )
266
+ def _handle_extract_sections(task, arguments, previous_pipelines):
267
+ """Reduce generated text to a known set of labelled sections"""
268
+ logger.debug("Extracting sections")
269
+ from .text_sections import extract_sections
270
+
271
+ return extract_sections(**arguments)
272
+
273
+
274
+ @register_command(
275
+ "batch_decode_post_process",
276
+ implementation="dw.tasks.format_messages.batch_decode_post_process",
277
+ provided=("processor",),
278
+ )
279
+ def _handle_batch_decode(task, arguments, previous_pipelines):
280
+ """Batch decode post-processing with pipeline reference"""
281
+ logger.debug("Performing batch decode post-processing")
282
+ pipeline_reference = task.task_definition["pipeline_reference"]
283
+ if pipeline_reference not in previous_pipelines:
284
+ raise KeyError(
285
+ f"Pipeline reference '{pipeline_reference}' not found in previous pipelines. "
286
+ f"Available pipelines: {list(previous_pipelines.keys())}"
287
+ )
288
+ processor = previous_pipelines[pipeline_reference].pipeline
289
+ return batch_decode_post_process(processor, **arguments)
290
+
291
+
292
+ def _handle_image_processing(task, arguments, previous_pipelines):
293
+ """Handle image processing commands"""
294
+ logger.debug("Processing image")
295
+ device = task.device_for(arguments)
296
+ return process_image(
297
+ arguments.pop("image"),
298
+ task.command,
299
+ device,
300
+ arguments,
301
+ )
302
+
303
+
304
+ def _handle_video_processing(task, arguments, previous_pipelines):
305
+ """Handle video processing commands"""
306
+ logger.debug("Processing video")
307
+ device = task.device_for(arguments)
308
+ return process_video(
309
+ arguments.pop("video"),
310
+ task.command,
311
+ device,
312
+ arguments,
313
+ )
314
+
315
+
316
+ # Command names process_video (video_utils.py) accepts, with the function
317
+ # whose signature carries their arguments. video_utils dispatches via a plain
318
+ # if-chain, so keep this in sync with the branches in process_video().
319
+ # get_first/last_frame pin frame_index themselves, so it is 'provided'.
320
+ _VIDEO_PROCESSOR_INFO = {
321
+ "get_frame": {
322
+ "kind": "video_processor",
323
+ "implementation": "dw.tasks.video_utils.get_frame",
324
+ "provided": (),
325
+ },
326
+ "get_first_frame": {
327
+ "kind": "video_processor",
328
+ "implementation": "dw.tasks.video_utils.get_frame",
329
+ "provided": ("frame_index",),
330
+ },
331
+ "get_last_frame": {
332
+ "kind": "video_processor",
333
+ "implementation": "dw.tasks.video_utils.get_frame",
334
+ "provided": ("frame_index",),
335
+ },
336
+ }
337
+ _VIDEO_PROCESSOR_COMMANDS = sorted(_VIDEO_PROCESSOR_INFO)
338
+
339
+
340
+ class Task:
341
+ """
342
+ Represents a task that can be executed as part of a workflow.
343
+ Tasks are atomic operations like image processing, data gathering, or message formatting.
344
+ """
345
+
346
+ def __init__(self, task_definition, device):
347
+ """
348
+ Initialize task with its configuration and device settings.
349
+
350
+ Args:
351
+ task_definition: Dictionary containing task configuration and parameters
352
+ device: Device to run task on (e.g., 'cuda', 'mps', 'cpu')
353
+ """
354
+ self.task_definition = task_definition
355
+ self.device = device
356
+ logger.debug(f"Initialized task: {self.name} for device: {device}")
357
+
358
+ @property
359
+ def name(self):
360
+ """Get task name from command property"""
361
+ return self.command
362
+
363
+ def device_for(self, arguments):
364
+ """Get the device this task runs on, consuming any override in its arguments.
365
+
366
+ A task can pin itself to a device - a captioning model on the CPU while the GPU
367
+ holds a pipeline, for instance. The argument is removed either way so it does
368
+ not reach the command as a duplicate.
369
+
370
+ Args:
371
+ arguments: Arguments for this run of the task
372
+
373
+ Returns:
374
+ Device identifier the task should run on
375
+ """
376
+ return arguments.pop("device", self.device)
377
+
378
+ @property
379
+ def argument_template(self):
380
+ """
381
+ Get argument template for this task.
382
+
383
+ Returns:
384
+ Dictionary of arguments from inputs or arguments section
385
+ """
386
+ # A task will either be an input array or a dictionary of arguments
387
+ if "inputs" in self.task_definition:
388
+ logger.debug("Using inputs as argument template")
389
+ return self.task_definition["inputs"]
390
+
391
+ logger.debug("Using arguments as argument template")
392
+ return self.task_definition["arguments"]
393
+
394
+ @property
395
+ def command(self):
396
+ """Get command name or 'unknown' if not specified"""
397
+ return self.task_definition.get("command", "unknown")
398
+
399
+ def run(self, arguments, previous_pipelines={}):
400
+ """
401
+ Execute the task with given arguments using the command registry.
402
+
403
+ Args:
404
+ arguments: Dictionary of arguments for task execution
405
+ previous_pipelines: Dictionary of previously created pipelines
406
+
407
+ Returns:
408
+ Task output based on command type
409
+
410
+ Raises:
411
+ ValueError: If command is unknown
412
+ KeyError: If required arguments or pipeline references are missing
413
+ """
414
+ logger.debug(f"Running task: {self.command}")
415
+ logger.debug(f"Task arguments: {arguments}")
416
+
417
+ try:
418
+ # Cooperative cancellation reaches task steps too - without this
419
+ # a cancel during a long task waits for the whole task to finish
420
+ from ..events import get_context
421
+
422
+ get_context().check_cancelled()
423
+
424
+ # Look up command in registry
425
+ if self.command in _COMMAND_REGISTRY:
426
+ handler = _COMMAND_REGISTRY[self.command]
427
+ return handler(self, arguments, previous_pipelines)
428
+
429
+ # Not a registered command - check whether it names an image or
430
+ # video processor instead. Imported lazily here to preserve
431
+ # image_utils' lazy-import discipline for callers that never
432
+ # touch image processing.
433
+ from .image_utils import available_processors
434
+
435
+ if self.command in available_processors():
436
+ return _handle_image_processing(self, arguments, previous_pipelines)
437
+
438
+ if self.command in _VIDEO_PROCESSOR_COMMANDS:
439
+ return _handle_video_processing(self, arguments, previous_pipelines)
440
+
441
+ # Unknown command - not in the registry, and not a known image or
442
+ # video processor name either
443
+ error_msg = (
444
+ f"Unknown task command: '{self.command}'. "
445
+ f"Registered commands: {sorted(_COMMAND_REGISTRY.keys())}. "
446
+ f"Image processors: {available_processors()}. "
447
+ f"Video processors: {_VIDEO_PROCESSOR_COMMANDS}"
448
+ )
449
+ logger.error(error_msg)
450
+ raise ValueError(error_msg)
451
+
452
+ except KeyError as e:
453
+ # Missing required arguments or pipeline references
454
+ logger.error(
455
+ f"Missing required data for task {self.command}: {e}", exc_info=True
456
+ )
457
+ raise
458
+ except (ValueError, TypeError) as e:
459
+ # Invalid arguments or type mismatches
460
+ logger.error(
461
+ f"Invalid arguments for task {self.command}: {e}", exc_info=True
462
+ )
463
+ raise
464
+ except (OSError, IOError) as e:
465
+ # File operations, resource loading errors
466
+ logger.error(f"I/O error in task {self.command}: {e}", exc_info=True)
467
+ raise
468
+ except Exception as e:
469
+ # Catch-all for unexpected errors
470
+ logger.error(
471
+ f"Unexpected error ({type(e).__name__}) executing task {self.command}: {e}",
472
+ exc_info=True,
473
+ )
474
+ raise
@@ -0,0 +1,57 @@
1
+ """
2
+ Shared PIL <-> float tensor conversions.
3
+
4
+ Consolidates the PIL-to-tensor round trip that was previously hand-rolled
5
+ independently in upscale.py, restore_faces.py, and interpolate_frames.py.
6
+ """
7
+
8
+ import numpy as np
9
+ import torch
10
+ from PIL import Image
11
+
12
+
13
+ def pil_to_float_tensor(image, device, dtype=None):
14
+ """Convert a PIL image to a (1, 3, H, W) float tensor in [0, 1] on device.
15
+
16
+ The image is coerced to RGB first, so single-channel or RGBA inputs are
17
+ handled consistently. `dtype` defaults to float32 (the numpy source
18
+ precision); pass e.g. `torch.float16` to cast directly to a model's
19
+ working precision.
20
+
21
+ Args:
22
+ image: PIL Image
23
+ device: Target device (str or torch.device)
24
+ dtype: Optional torch dtype to cast to (default: float32)
25
+
26
+ Returns:
27
+ torch.Tensor of shape (1, 3, H, W)
28
+ """
29
+ arr = np.array(image.convert("RGB")).astype(np.float32) / 255.0
30
+ tensor = torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0).to(device)
31
+ if dtype is not None:
32
+ tensor = tensor.to(dtype=dtype)
33
+ return tensor
34
+
35
+
36
+ def float_tensor_to_pil(tensor):
37
+ """Convert a (1, 3, H, W) or (3, H, W) float tensor in [0, 1] to a PIL RGB image.
38
+
39
+ Quantizes to uint8 via rounding (`.round()`) rather than truncation
40
+ (a bare truncating cast), matching diffusers' VaeImageProcessor.numpy_to_pil
41
+ behavior. This matters: with truncation, an exact 8-bit value that
42
+ round-trips through [0, 1] float (e.g. 128/255) can land a hair below
43
+ its integer (127.999...) and get chopped down a level instead of
44
+ landing back on 128. Rounding fixes that, at the cost of at most 1/255
45
+ of drift per channel versus the old truncating behavior for values that
46
+ were never exact to begin with.
47
+
48
+ Args:
49
+ tensor: torch.Tensor of shape (1, 3, H, W) or (3, H, W), values in [0, 1]
50
+
51
+ Returns:
52
+ PIL.Image.Image in RGB mode
53
+ """
54
+ if tensor.dim() == 4:
55
+ tensor = tensor.squeeze(0)
56
+ arr = tensor.permute(1, 2, 0).mul(255).round().clamp(0, 255).byte().cpu().numpy()
57
+ return Image.fromarray(arr)
@@ -0,0 +1,168 @@
1
+ """
2
+ Text generation via HuggingFace transformers.
3
+
4
+ Takes a prompt (and optional system prompt) and generates text using a
5
+ local language model. Useful for prompt expansion, rewriting, and
6
+ other text-to-text tasks.
7
+
8
+ Supplying an image switches to a vision-language model, so the generated
9
+ text can describe what is actually in the picture rather than what the
10
+ prompt guesses is there. That is the difference between an image-conditioned
11
+ workflow whose prompt agrees with its keyframe and one whose prompt fights it.
12
+ """
13
+
14
+ import logging
15
+ from transformers import pipeline as hf_pipeline
16
+ from .. import preferred_task_dtype
17
+ from .model_cache import cached_model
18
+
19
+ logger = logging.getLogger("dw")
20
+
21
+ _DEFAULT_MODEL = "Qwen/Qwen2.5-1.5B-Instruct"
22
+ # Small enough to stand in for the old captioning default; override for
23
+ # anything needing real detail
24
+ _DEFAULT_VISION_MODEL = "HuggingFaceTB/SmolVLM-256M-Instruct"
25
+
26
+ # Greedy decoding against a long, rigid format specification makes the vision
27
+ # models loop - finishing the answer, then repeating its closing sections until
28
+ # they run out of tokens. A penalty stops that without giving up reproducible
29
+ # output, which sampling would. Measured on Qwen3-VL against the H3 prompt
30
+ # spec: 1.05 still looped and filled the whole budget, 1.15 ended on its own at
31
+ # a length matching the format's own guidance. The text models do not need it
32
+ _VISION_REPETITION_PENALTY = 1.15
33
+
34
+
35
+ def _image_part(image):
36
+ """The chat content entry for an image.
37
+
38
+ A PIL image goes in as an object; a string is a URL or path the pipeline
39
+ loads itself, and it has to be declared as such - passing it under "image"
40
+ would hand the processor a bare string where it expects pixels.
41
+ """
42
+ return {"type": "image", "url" if isinstance(image, str) else "image": image}
43
+
44
+
45
+ def _build_messages(prompt, system_prompt, image):
46
+ """Chat messages in the shape the chosen pipeline expects.
47
+
48
+ Text-generation models take plain string content. Vision models take a
49
+ list of typed parts, because the image is a part of the message rather
50
+ than something alongside it.
51
+ """
52
+ messages = []
53
+ if image is None:
54
+ if system_prompt is not None:
55
+ messages.append({"role": "system", "content": system_prompt})
56
+ messages.append({"role": "user", "content": prompt})
57
+ else:
58
+ if system_prompt is not None:
59
+ messages.append(
60
+ {"role": "system", "content": [{"type": "text", "text": system_prompt}]}
61
+ )
62
+ messages.append(
63
+ {
64
+ "role": "user",
65
+ "content": [_image_part(image), {"type": "text", "text": prompt}],
66
+ }
67
+ )
68
+ return messages
69
+
70
+
71
+ def generate_text(prompt, device="cpu", **kwargs):
72
+ """Generate text from a prompt using a local language model.
73
+
74
+ Args:
75
+ prompt: The user message / prompt to expand or transform.
76
+ device: Target device ("cuda", "mps", "cpu").
77
+ **kwargs:
78
+ model_name: HuggingFace model ID. Defaults to
79
+ Qwen/Qwen2.5-1.5B-Instruct, or a small vision-language model
80
+ when an image is supplied. An image needs a model that can
81
+ accept one - a text-only model will fail to load as one.
82
+ system_prompt: Optional system instruction for the model.
83
+ max_new_tokens: Max tokens to generate (default: 500).
84
+ image: Optional PIL image, URL or path. Its presence is what
85
+ selects the vision pipeline.
86
+ repetition_penalty: Vision pipeline only (default: 1.15). Raise it
87
+ if a model still repeats itself, or set 1.0 to disable.
88
+ generate_kwargs: Anything else to hand the model's generate() -
89
+ no_repeat_ngram_size, top_p, min_new_tokens and so on. Merged
90
+ over what this function sets, so it can override those too.
91
+
92
+ Returns:
93
+ Generated text string.
94
+ """
95
+ # A workflow declaring an optional image passes it through as null when the
96
+ # caller supplies none, and an empty string is the same statement
97
+ image = kwargs.get("image", None) or None
98
+ system_prompt = kwargs.get("system_prompt", None)
99
+ max_new_tokens = int(kwargs.get("max_new_tokens", 500))
100
+
101
+ if image is None:
102
+ pipeline_task = "text-generation"
103
+ model_name = kwargs.get("model_name", _DEFAULT_MODEL)
104
+ else:
105
+ pipeline_task = "image-text-to-text"
106
+ model_name = kwargs.get("model_name", _DEFAULT_VISION_MODEL)
107
+
108
+ dtype = preferred_task_dtype(device)
109
+
110
+ def load_pipe():
111
+ logger.info(f"Generating text with {model_name} on {device}")
112
+ return hf_pipeline(
113
+ pipeline_task,
114
+ model=model_name,
115
+ device_map=device,
116
+ torch_dtype=dtype,
117
+ )
118
+
119
+ # The task is part of the identity - the same model name can be loaded
120
+ # under either pipeline, and they are not interchangeable
121
+ pipe = cached_model(
122
+ ("text_generation", pipeline_task, model_name, str(device), str(dtype)),
123
+ load_pipe,
124
+ )
125
+
126
+ messages = _build_messages(prompt, system_prompt, image)
127
+
128
+ # Decoding is greedy, so the same input returns the same text every run -
129
+ # which is what a workflow wants, and why there is nothing here to seed
130
+ generation = {"do_sample": False}
131
+ if image is not None:
132
+ generation["repetition_penalty"] = float(
133
+ kwargs.get("repetition_penalty", _VISION_REPETITION_PENALTY)
134
+ )
135
+ # Last, so a workflow can override anything decided above
136
+ generation.update(kwargs.get("generate_kwargs") or {})
137
+
138
+ return _generate(pipe, messages, image, max_new_tokens, generation)
139
+
140
+
141
+ def _generate(pipe, messages, image, max_new_tokens, generation):
142
+ """Run the pipeline, which each path calls differently."""
143
+ if image is None:
144
+ # This pipeline collects anything it does not name into the arguments it
145
+ # forwards to generate(), so settings go in as plain keywords
146
+ results = pipe(
147
+ messages,
148
+ max_new_tokens=max_new_tokens,
149
+ return_full_text=False,
150
+ **generation,
151
+ )
152
+ else:
153
+ # Images live inside the messages, so the chat goes in as `text` - the
154
+ # pipeline rejects a chat and an `images` argument together. Generation
155
+ # settings have to go through generate_kwargs here: anything else this
156
+ # pipeline does not name explicitly is forwarded to the processor and
157
+ # dropped, so a bare do_sample=False would leave sampling on. Passing
158
+ # max_new_tokens both ways is an error, so it stays a direct argument
159
+ results = pipe(
160
+ text=messages,
161
+ max_new_tokens=max_new_tokens,
162
+ return_full_text=False,
163
+ generate_kwargs=generation,
164
+ )
165
+
166
+ text = results[0]["generated_text"].strip()
167
+ logger.info(f"Generated: {text[:100]}{'...' if len(text) > 100 else ''}")
168
+ return text