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/result.py ADDED
@@ -0,0 +1,850 @@
1
+ import os
2
+ import numpy
3
+ import torch
4
+ import soundfile
5
+ import json
6
+ import mimetypes
7
+ import logging
8
+ from diffusers.utils import (
9
+ export_to_video,
10
+ export_to_gif,
11
+ encode_video,
12
+ is_av_available,
13
+ )
14
+ from collections.abc import Mapping
15
+ from .security import validate_output_path, validate_string_input, SecurityError
16
+
17
+ logger = logging.getLogger("dw")
18
+
19
+ # Result saving constants
20
+ MAX_BASE_NAME_LENGTH = 200
21
+ DEFAULT_AUDIO_SAMPLE_RATE = 44100
22
+
23
+ # Audio content types soundfile can write, mapped to their file extension and to any
24
+ # write arguments the extension alone does not imply. Opus has no extension of its own
25
+ # in libsndfile - it is a subtype of the ogg container.
26
+ AUDIO_FORMATS = {
27
+ "audio/wav": (".wav", {}),
28
+ "audio/x-wav": (".wav", {}),
29
+ "audio/aiff": (".aiff", {}),
30
+ "audio/flac": (".flac", {}),
31
+ "audio/x-flac": (".flac", {}),
32
+ "audio/mpeg": (".mp3", {}),
33
+ "audio/mp3": (".mp3", {}),
34
+ "audio/ogg": (".ogg", {}),
35
+ "audio/vorbis": (".ogg", {}),
36
+ "audio/opus": (".ogg", {"format": "OGG", "subtype": "OPUS"}),
37
+ }
38
+
39
+ # Result definition keys passed through to soundfile - encoding quality controls
40
+ AUDIO_WRITE_ARGUMENTS = ["subtype", "format", "compression_level", "bitrate_mode"]
41
+
42
+ # Audio is written in chunks of this many frames - see write_audio
43
+ AUDIO_WRITE_CHUNK_FRAMES = 1 << 20
44
+
45
+ # The only container encode_video writes - it always encodes h264 video
46
+ MUXED_VIDEO_CONTENT_TYPE = "video/mp4"
47
+
48
+ # Distinguishes "the artifact has no such attribute" from "it has one holding None" -
49
+ # an AudioVideo whose pipeline reported no sample rate carries exactly that
50
+ _NO_PROPERTY = object()
51
+
52
+ # The names a modular pipeline's outputs go by. Asked for more than one output it returns
53
+ # them in a dict rather than on a pipeline output object, so its videos and the soundtrack
54
+ # generated alongside them arrive keyed instead of as attributes. Every diffusers modular
55
+ # pipeline (minimax_h3, ltx2, ...) names these "videos"/"audio"/"sampling_rate" - kept as
56
+ # tuples, rather than plain strings, only so the lookup goes through first_item like the
57
+ # other key sets. "audio_sample_rate" is real too: it is the name dw's own
58
+ # attach_audio_sample_rate (pipeline_processors/pipeline.py) gives the rate when it
59
+ # attaches it to a non-modular output - a modular result carrying it under that name is
60
+ # tested (TestModularOutputs.test_audio_sample_rate_names_the_rate_too) and kept for it.
61
+ MODULAR_VIDEO_KEYS = ("videos",)
62
+ MODULAR_AUDIO_KEYS = ("audio",)
63
+ MODULAR_SAMPLE_RATE_KEYS = ("sampling_rate", "audio_sample_rate")
64
+
65
+
66
+ class AudioVideo:
67
+ """A generated video together with the audio track generated alongside it.
68
+
69
+ Pipelines like LTX-2 return audio next to their frames. Keeping the two paired lets
70
+ the result mux them into one file instead of dropping the audio on the floor.
71
+ """
72
+
73
+ def __init__(self, frames, audio, sample_rate):
74
+ """
75
+ Args:
76
+ frames: The video, as PIL images or an array of frames
77
+ audio: Waveform for this video, shaped (channels, samples)
78
+ sample_rate: Sample rate of the waveform, or None if the pipeline did not report one
79
+ """
80
+ self.frames = frames
81
+ self.audio = audio
82
+ self.sample_rate = sample_rate
83
+
84
+
85
+ class Result:
86
+ """Manages and stores results from workflow steps.
87
+
88
+ Handles result storage, artifact management, and file saving with support
89
+ for multiple content types including images, video, audio, and JSON.
90
+ """
91
+
92
+ def __init__(self, result_definition):
93
+ """Initialize Result with configuration for how to handle/save results.
94
+
95
+ Args:
96
+ result_definition: Dict containing result configuration including:
97
+ - content_type: MIME type of the result
98
+ - save: Boolean indicating if result should be saved
99
+ - file_base_name: Base name for saved files
100
+ """
101
+ self.result_definition = result_definition
102
+ self.result_list = []
103
+ self.metadata = None
104
+ self.saved_files = []
105
+ logger.debug(f"Initialized Result with definition: {result_definition}")
106
+
107
+ def set_metadata(self, metadata):
108
+ """Set metadata to embed in saved image artifacts.
109
+
110
+ Args:
111
+ metadata: Dict of generation parameters to embed
112
+ """
113
+ self.metadata = metadata
114
+
115
+ def add_result(self, result):
116
+ """Add one or more results to the result list.
117
+
118
+ Args:
119
+ result: Single result or list of results to store
120
+ """
121
+ if isinstance(result, list):
122
+ logger.debug(f"Adding {len(result)} results to result list")
123
+ self.result_list.extend(result)
124
+ else:
125
+ if isinstance(result, str):
126
+ # Clean up string results by removing extra quotes and whitespace
127
+ result = result.strip().strip('"').strip()
128
+ logger.debug("Adding single result to result list")
129
+ self.result_list.append(result)
130
+
131
+ def get_artifacts(self):
132
+ """Retrieve all artifacts from stored results.
133
+
134
+ Returns:
135
+ List of all artifacts from all results
136
+ """
137
+ artifacts = []
138
+ for result in self.result_list:
139
+ artifacts.extend(get_artifact_list(result))
140
+
141
+ logger.debug(f"Retrieved {len(artifacts)} artifacts from results")
142
+ return artifacts
143
+
144
+ def get_artifact_properties(self, property_name):
145
+ """Extract specific properties from results.
146
+
147
+ A dict result is looked up by key; anything else by attribute, which is how
148
+ a step reaches into an artifact that is an object rather than a mapping -
149
+ the frames or the soundtrack of the AudioVideo a video-with-audio pipeline
150
+ produces, say, where the next step takes one of them on its own. Methods are
151
+ not properties: 'previous_result:step.index' on a list result names nothing
152
+ the workflow meant, so it fails rather than passing a bound method along.
153
+
154
+ Args:
155
+ property_name: Name of property to extract from results
156
+
157
+ Returns:
158
+ List of property values from results where property exists
159
+
160
+ Raises:
161
+ ValueError: If a result is neither a dict-like (Mapping) object with that
162
+ key nor an object carrying it as a data attribute - a plain string or
163
+ other scalar result has no properties to look up, and staying quiet
164
+ about that (or doing a membership/substring test instead of a key
165
+ lookup) would silently drop data or raise a confusing TypeError.
166
+ """
167
+ values = []
168
+ for result in self.result_list:
169
+ if isinstance(result, Mapping):
170
+ if property_name in result:
171
+ values.append(result[property_name])
172
+ continue
173
+
174
+ value = getattr(result, property_name, _NO_PROPERTY)
175
+ # A string's every 'property' is a method, and so is most of a list's -
176
+ # the original loud failure for those is the useful answer
177
+ if value is _NO_PROPERTY or callable(value):
178
+ raise ValueError(
179
+ f"result has no property '{property_name}' "
180
+ f"(it is a {type(result).__name__}, not a dict)"
181
+ )
182
+ values.append(value)
183
+
184
+ logger.debug(f"Retrieved {len(values)} values for property: {property_name}")
185
+ return values
186
+
187
+ def save(self, output_dir, default_base_name):
188
+ """Save results to files based on content type.
189
+
190
+ Args:
191
+ output_dir: Directory to save files in
192
+ default_base_name: Default name to use for files
193
+
194
+ Returns:
195
+ List of file paths written, in the order they were written - the
196
+ step's manifest. Empty when saving is disabled or nothing saved.
197
+ """
198
+ try:
199
+ # Validate output directory
200
+ validated_output_dir = validate_output_path(output_dir, None)
201
+ validated_base_name = validate_string_input(
202
+ default_base_name, max_length=MAX_BASE_NAME_LENGTH
203
+ )
204
+
205
+ # Add directory check/creation
206
+ if not os.path.exists(validated_output_dir):
207
+ logger.debug(f"Creating output directory: {validated_output_dir}")
208
+ os.makedirs(validated_output_dir, exist_ok=True)
209
+ elif not os.path.isdir(validated_output_dir):
210
+ raise ValueError(
211
+ f"Output path exists but is not a directory: {validated_output_dir}"
212
+ )
213
+ except SecurityError as e:
214
+ logger.error(f"Security validation failed for output: {e}")
215
+ raise
216
+ except (OSError, PermissionError) as e:
217
+ logger.error(f"Failed to create output directory: {e}")
218
+ raise
219
+
220
+ # Check if saving is enabled and content type is specified
221
+ content_type = self.result_definition.get("content_type", None)
222
+ if not self.result_definition.get("save", True) or content_type is None:
223
+ logger.debug("Skipping save - disabled or no content type specified")
224
+ self.saved_files = []
225
+ return self.saved_files
226
+
227
+ # Determine base filename with validation
228
+ file_base_name = validated_base_name
229
+ if "file_base_name" in self.result_definition:
230
+ custom_base = validate_string_input(
231
+ self.result_definition["file_base_name"],
232
+ max_length=MAX_BASE_NAME_LENGTH,
233
+ )
234
+ file_base_name = custom_base + validated_base_name
235
+
236
+ # Get file extension for content type
237
+ extension = guess_extension(content_type)
238
+ logger.debug(
239
+ f"Saving with content type: {content_type}, extension: {extension}"
240
+ )
241
+
242
+ # Save each result, collecting the paths written as the step's manifest
243
+ saved_files = []
244
+ for i, result in enumerate(self.result_list):
245
+ if content_type.endswith("json"):
246
+ # Handle JSON content type
247
+ output_path = os.path.join(
248
+ validated_output_dir, f"{file_base_name}-{i}{extension}"
249
+ )
250
+ logger.info(f"Saving JSON result to {output_path}")
251
+ with open(output_path, "w") as file:
252
+ file.write(json.dumps(result, indent=4))
253
+ saved_files.append(output_path)
254
+ else:
255
+ # Handle other content types
256
+ for j, artifact in enumerate(get_artifact_list(result)):
257
+ saved_files.extend(
258
+ self.save_artifact(
259
+ validated_output_dir,
260
+ artifact,
261
+ f"{file_base_name}-{i}.{j}",
262
+ content_type,
263
+ extension,
264
+ )
265
+ )
266
+ self.saved_files = saved_files
267
+ return saved_files
268
+
269
+ def save_artifact(
270
+ self, output_dir, artifact, file_base_name, content_type, extension
271
+ ):
272
+ """Save individual artifact to file based on its type.
273
+
274
+ Args:
275
+ output_dir: Directory to save file in, already validated by save()
276
+ artifact: The artifact to save
277
+ file_base_name: Base name for the file, derived from names save()
278
+ validated
279
+ content_type: MIME type of the content
280
+ extension: File extension to use
281
+
282
+ Returns:
283
+ List of file paths written.
284
+ """
285
+ if artifact is None:
286
+ logger.warning(f"Skipping None artifact for {file_base_name}")
287
+ return []
288
+
289
+ if isinstance(artifact, dict):
290
+ # Recursively save dictionary items
291
+ logger.debug(
292
+ f"Saving dictionary artifact with keys: {list(artifact.keys())}"
293
+ )
294
+ saved_files = []
295
+ for k, v in artifact.items():
296
+ saved_files.extend(
297
+ self.save_artifact(
298
+ output_dir,
299
+ v,
300
+ f"{file_base_name}-{k}",
301
+ content_type,
302
+ extension,
303
+ )
304
+ )
305
+ return saved_files
306
+
307
+ output_path = os.path.join(output_dir, f"{file_base_name}{extension}")
308
+ logger.info(f"Saving artifact to {output_path}")
309
+
310
+ try:
311
+ if content_type.startswith("video"):
312
+ if isinstance(artifact, AudioVideo):
313
+ self.save_audio_video(artifact, output_path, content_type)
314
+ else:
315
+ export_to_video(
316
+ artifact, output_path, fps=self.result_definition.get("fps", 8)
317
+ )
318
+ elif content_type == "image/gif":
319
+ export_to_gif(
320
+ artifact, output_path, fps=self.result_definition.get("fps", 8)
321
+ )
322
+ elif content_type.startswith("audio"):
323
+ waveforms = normalize_audio(artifact)
324
+ sample_rate = self.result_definition.get(
325
+ "sample_rate",
326
+ self.result_definition.get("samplerate", DEFAULT_AUDIO_SAMPLE_RATE),
327
+ )
328
+ # A batched waveform holds several songs - save each one separately
329
+ if len(waveforms) > 1:
330
+ saved_files = []
331
+ for k, waveform in enumerate(waveforms):
332
+ saved_files.extend(
333
+ self.save_artifact(
334
+ output_dir,
335
+ waveform,
336
+ f"{file_base_name}-{k}",
337
+ content_type,
338
+ extension,
339
+ )
340
+ )
341
+ return saved_files
342
+ write_audio(
343
+ output_path,
344
+ waveforms[0],
345
+ sample_rate,
346
+ **self.get_audio_write_arguments(content_type),
347
+ )
348
+ elif content_type.endswith("json"):
349
+ with open(output_path, "w") as file:
350
+ file.write(json.dumps(artifact, indent=4))
351
+ elif content_type.startswith("text"):
352
+ with open(output_path, "w") as file:
353
+ file.write(artifact)
354
+ elif hasattr(artifact, "save"):
355
+ if (
356
+ self.metadata is not None
357
+ and self.result_definition.get("embed_metadata", False)
358
+ and content_type.startswith("image/")
359
+ ):
360
+ self._save_image_with_metadata(artifact, output_path, content_type)
361
+ else:
362
+ artifact.save(output_path)
363
+ else:
364
+ raise ValueError(
365
+ f"Content type {content_type} does not match result type {type(artifact)}"
366
+ )
367
+ except Exception as e:
368
+ logger.error(
369
+ f"Error saving artifact to {output_path}: {str(e)}", exc_info=True
370
+ )
371
+ raise
372
+
373
+ return [output_path]
374
+
375
+ def save_audio_video(self, artifact, output_path, content_type):
376
+ """Write a video and the audio generated with it into a single file.
377
+
378
+ encode_video muxes the two into an h264/mp4 file with PyAV. When PyAV is missing,
379
+ the container is not mp4, or nothing told us the sample rate, the video is written
380
+ on its own and the audio is dropped.
381
+
382
+ Args:
383
+ artifact: AudioVideo holding the frames and their waveform
384
+ output_path: Path of the file to write
385
+ content_type: MIME type of the video being written
386
+ """
387
+ fps = self.result_definition.get("fps", 8)
388
+ # The pipeline reports the sample rate of what it generated - the result
389
+ # definition can still override it
390
+ sample_rate = self.result_definition.get(
391
+ "audio_sample_rate", artifact.sample_rate
392
+ )
393
+
394
+ # Segment-backed frames (a chained step with save_segments) replay from
395
+ # disk one segment at a time, so the final video is streamed instead of
396
+ # materialized - and the segment files are removed once it is written
397
+ if hasattr(artifact.frames, "cleanup"):
398
+ if content_type != MUXED_VIDEO_CONTENT_TYPE:
399
+ raise ValueError(
400
+ f"Segment-backed video can only be written as "
401
+ f"{MUXED_VIDEO_CONTENT_TYPE}, not {content_type}"
402
+ )
403
+ if not is_av_available():
404
+ raise ValueError(
405
+ "Writing segment-backed video needs PyAV - install it "
406
+ "with: pip install av"
407
+ )
408
+
409
+ audio = None
410
+ if artifact.audio is not None and sample_rate is not None:
411
+ audio = as_audio_track(artifact.audio)
412
+ logger.debug(
413
+ f"Streaming {len(artifact.frames)} segments into {output_path}"
414
+ )
415
+ encode_video(
416
+ iter(artifact.frames),
417
+ fps=fps,
418
+ output_path=output_path,
419
+ audio=audio,
420
+ audio_sample_rate=sample_rate if audio is not None else None,
421
+ video_chunks_number=len(artifact.frames),
422
+ )
423
+ artifact.frames.cleanup()
424
+ return
425
+
426
+ reason = None
427
+ if artifact.audio is None:
428
+ reason = "the pipeline returned no audio"
429
+ elif sample_rate is None:
430
+ reason = "the audio sample rate is unknown"
431
+ elif content_type != MUXED_VIDEO_CONTENT_TYPE:
432
+ reason = f"audio can only be muxed into {MUXED_VIDEO_CONTENT_TYPE}"
433
+ elif not is_av_available():
434
+ reason = "PyAV is not installed - install it with: pip install av"
435
+
436
+ if reason is not None:
437
+ # No audio at all is an expected shape - video-only chains and
438
+ # concatenations - so it logs quietly; losing audio we do have warns
439
+ log = logger.debug if artifact.audio is None else logger.warning
440
+ log(f"Saving {output_path} without its audio because {reason}")
441
+ export_to_video(artifact.frames, output_path, fps=fps)
442
+ return
443
+
444
+ logger.debug(f"Muxing audio at {sample_rate}Hz into {output_path}")
445
+ encode_video(
446
+ artifact.frames,
447
+ fps=fps,
448
+ output_path=output_path,
449
+ audio=as_audio_track(artifact.audio),
450
+ audio_sample_rate=sample_rate,
451
+ )
452
+
453
+ def get_audio_write_arguments(self, content_type):
454
+ """Collect the soundfile arguments for an audio content type.
455
+
456
+ The container comes from the content type and any encoding quality settings
457
+ from the result definition.
458
+
459
+ Args:
460
+ content_type: MIME type of the audio being written
461
+
462
+ Returns:
463
+ Dict of keyword arguments for soundfile.write
464
+ """
465
+ _, write_arguments = AUDIO_FORMATS.get(content_type, (None, {}))
466
+ write_arguments = dict(write_arguments)
467
+
468
+ for argument_name in AUDIO_WRITE_ARGUMENTS:
469
+ value = self.result_definition.get(argument_name, None)
470
+ if value is not None:
471
+ write_arguments[argument_name] = value
472
+
473
+ return write_arguments
474
+
475
+ def _save_image_with_metadata(self, image, output_path, content_type):
476
+ """Save an image with embedded generation metadata.
477
+
478
+ Args:
479
+ image: PIL Image to save
480
+ output_path: Path to save the image to
481
+ content_type: MIME type of the image
482
+ """
483
+ metadata_json = json.dumps(self.metadata, default=str)
484
+
485
+ if content_type == "image/png":
486
+ from PIL.PngImagePlugin import PngInfo
487
+
488
+ png_info = PngInfo()
489
+ png_info.add_text("parameters", metadata_json)
490
+ image.save(output_path, pnginfo=png_info)
491
+ logger.debug(f"Embedded PNG metadata in {output_path}")
492
+ elif content_type in ("image/jpeg", "image/webp"):
493
+ try:
494
+ import piexif
495
+ import piexif.helper
496
+
497
+ exif_dict = {"0th": {}, "Exif": {}, "GPS": {}, "1st": {}}
498
+ if hasattr(image, "info") and "exif" in image.info:
499
+ exif_dict = piexif.load(image.info["exif"])
500
+ exif_dict["Exif"][piexif.ExifIFD.UserComment] = (
501
+ piexif.helper.UserComment.dump(metadata_json)
502
+ )
503
+ exif_bytes = piexif.dump(exif_dict)
504
+ image.save(output_path, exif=exif_bytes)
505
+ logger.debug(f"Embedded EXIF metadata in {output_path}")
506
+ except ImportError:
507
+ logger.warning(
508
+ "piexif not installed - saving without metadata. "
509
+ "Install with: pip install piexif"
510
+ )
511
+ image.save(output_path)
512
+ else:
513
+ image.save(output_path)
514
+
515
+
516
+ def read_embedded_metadata(path):
517
+ """The generation metadata a saved image carries, or None.
518
+
519
+ The read-side mirror of _save_image_with_metadata: the 'parameters' PNG
520
+ text chunk, or the EXIF UserComment for JPEG/WebP. Returns the parsed
521
+ dict, or None when the file has no metadata this writer produced.
522
+ """
523
+ try:
524
+ from PIL import Image
525
+
526
+ with Image.open(path) as image:
527
+ text = getattr(image, "text", {}).get("parameters")
528
+ if text is None and "exif" in getattr(image, "info", {}):
529
+ import piexif
530
+ import piexif.helper
531
+
532
+ exif = piexif.load(image.info["exif"])
533
+ comment = exif.get("Exif", {}).get(piexif.ExifIFD.UserComment)
534
+ if comment:
535
+ text = piexif.helper.UserComment.load(comment)
536
+ if text is None:
537
+ return None
538
+ parsed = json.loads(text)
539
+ return parsed if isinstance(parsed, dict) else None
540
+ except Exception as e:
541
+ logger.debug(f"No readable metadata in {path}: {e}")
542
+ return None
543
+
544
+
545
+ def _frames_from_attributes(result):
546
+ """The frames extractor for a pipeline output that carries `.frames` directly.
547
+
548
+ Some video pipelines (LTX-2) generate an audio track along with the frames, exposed
549
+ as `.audio` and `.audio_sample_rate` attributes alongside `.frames`.
550
+ """
551
+ return frames_with_audio(
552
+ result.frames,
553
+ getattr(result, "audio", None),
554
+ getattr(result, "audio_sample_rate", None),
555
+ )
556
+
557
+
558
+ def _audios_from_attribute(result):
559
+ return [as_waveform_array(audio) for audio in result.audios]
560
+
561
+
562
+ # Diffusers output fields get_artifact_list knows how to turn into artifacts, tried in
563
+ # this order. A result is dispatched to the first field it has - images wins over frames
564
+ # if a result somehow has both, matching the fixed hasattr chain this replaced. Supporting
565
+ # a new diffusers output field (e.g. a standalone "depth" attribute) is one more entry
566
+ # here, instead of another branch threaded through the chain.
567
+ OUTPUT_FIELD_EXTRACTORS = [
568
+ ("images", lambda result: result.images),
569
+ ("image_embeds", lambda result: result.image_embeds),
570
+ ("image_embeddings", lambda result: result.image_embeddings),
571
+ ("frames", _frames_from_attributes),
572
+ ("audios", _audios_from_attribute),
573
+ ]
574
+
575
+
576
+ def get_artifact_list(result):
577
+ """Extract list of artifacts from a result object.
578
+
579
+ Handles various result types including images, embeddings, frames, and audio.
580
+
581
+ Args:
582
+ result: Result object to extract artifacts from
583
+
584
+ Returns:
585
+ List of artifacts
586
+ """
587
+ # Already a paired video and audio track - it has a .frames attribute of its own,
588
+ # but the pair is one artifact, not something to run back through frame extraction
589
+ if isinstance(result, AudioVideo):
590
+ return [result]
591
+
592
+ for field_name, extract in OUTPUT_FIELD_EXTRACTORS:
593
+ if hasattr(result, field_name):
594
+ return extract(result)
595
+
596
+ if isinstance(result, dict):
597
+ # A modular pipeline asked for several outputs returns them keyed
598
+ artifacts = modular_artifacts(result)
599
+ if artifacts is not None:
600
+ return artifacts
601
+
602
+ if isinstance(result, list):
603
+ return result
604
+
605
+ if hasattr(result, "to_tuple") or hasattr(result, "__dataclass_fields__"):
606
+ # A diffusers output (BaseOutput subclasses have to_tuple; plain dataclasses
607
+ # have __dataclass_fields__) whose fields matched none of the extractors above -
608
+ # log what it actually looks like so the resulting content-type mismatch on save
609
+ # is diagnosable instead of a bare "does not match result type" surprise.
610
+ logger.warning(
611
+ f"Don't know how to extract artifacts from a {type(result).__name__} - "
612
+ f"treating it as a single artifact. Its fields are: {output_field_names(result)}"
613
+ )
614
+
615
+ return [result]
616
+
617
+
618
+ def output_field_names(result):
619
+ """Best-effort list of field names on a diffusers-style output object, for logging."""
620
+ if hasattr(result, "keys"):
621
+ return list(result.keys())
622
+ if hasattr(result, "__dataclass_fields__"):
623
+ return list(result.__dataclass_fields__.keys())
624
+ return []
625
+
626
+
627
+ def modular_artifacts(result):
628
+ """Extract the artifacts from the outputs a modular pipeline returns together.
629
+
630
+ Asked for several outputs - `"output": ["videos", "audio", "sampling_rate"]` - a
631
+ modular pipeline returns them in a dict instead of on one output object. Pairing the
632
+ videos with the audio generated alongside them here saves them the same way a video
633
+ pipeline's own output is saved, muxed into a single file. Any other requested output
634
+ - "images" or "latents", say - is not part of that pairing, so it is carried along as
635
+ one extra dict artifact, saved key by key the same way any other dictionary result is.
636
+
637
+ Args:
638
+ result: Dict of outputs returned by a modular pipeline
639
+
640
+ Returns:
641
+ List of artifacts, or None when the outputs hold no video - those are saved one
642
+ output at a time instead
643
+ """
644
+ video_key, videos = first_item(result, MODULAR_VIDEO_KEYS)
645
+ if videos is None:
646
+ return None
647
+
648
+ consumed_keys = {video_key}
649
+
650
+ audio_key, audio = first_item(result, MODULAR_AUDIO_KEYS)
651
+ sample_rate = None
652
+ if audio is not None:
653
+ consumed_keys.add(audio_key)
654
+ rate_key, sample_rate = first_item(result, MODULAR_SAMPLE_RATE_KEYS)
655
+ consumed_keys.add(rate_key)
656
+
657
+ artifacts = frames_with_audio(videos, audio, sample_rate)
658
+
659
+ # Keys the video/audio pairing above did not consume still need to be saved, not
660
+ # dropped - carry them along as one extra artifact, saved key by key like any other
661
+ # dictionary result
662
+ leftovers = {
663
+ key: value
664
+ for key, value in result.items()
665
+ if key not in consumed_keys and value is not None
666
+ }
667
+ if leftovers:
668
+ artifacts = list(artifacts) + [leftovers]
669
+
670
+ return artifacts
671
+
672
+
673
+ def first_item(values, keys):
674
+ """The key and value of the first of `keys` present in `values`.
675
+
676
+ Returns (None, None) when none of them are.
677
+ """
678
+ for key in keys:
679
+ value = values.get(key, None)
680
+ if value is not None:
681
+ return key, value
682
+
683
+ return None, None
684
+
685
+
686
+ def frames_with_audio(frames, audio, sample_rate):
687
+ """Pair frames with the audio track generated alongside them, if there is one.
688
+
689
+ The one place that decides whether frames need pairing with audio at all - used by
690
+ both routes a pipeline's frames-plus-audio output can take: attributes on a pipeline
691
+ output object (`_frames_from_attributes`), and keys in the dict a modular pipeline
692
+ returns (`modular_artifacts`). Frames without audio are returned unchanged; actually
693
+ pairing them is `pair_audio_with_frames`'s job.
694
+
695
+ Args:
696
+ frames: The generated video(s), one list of frames per generation
697
+ audio: The generated waveform(s), or None if the pipeline produced no audio
698
+ sample_rate: Sample rate of the waveform(s), or None if unknown
699
+
700
+ Returns:
701
+ `frames` unchanged if `audio` is None, otherwise the list of AudioVideo pairs
702
+ `pair_audio_with_frames` produces
703
+ """
704
+ if audio is None:
705
+ return frames
706
+ return pair_audio_with_frames(frames, audio, sample_rate)
707
+
708
+
709
+ def pair_audio_with_frames(videos, audio, sample_rate):
710
+ """Pair each generated video with its own audio track.
711
+
712
+ Both are batched - videos[i] and audio[i] belong to the same generation. Only the
713
+ pipeline knows the sample rate its vocoder produced the audio at, so it comes along
714
+ rather than being guessed at here.
715
+
716
+ Args:
717
+ videos: The generated videos, one list of frames per generation
718
+ audio: The generated waveforms, one per generation
719
+ sample_rate: Sample rate of the waveforms, or None if the pipeline did not report one
720
+
721
+ Returns:
722
+ List of AudioVideo artifacts, one per generated video
723
+ """
724
+ return [
725
+ AudioVideo(frames, audio[i] if i < len(audio) else None, sample_rate)
726
+ for i, frames in enumerate(videos)
727
+ ]
728
+
729
+
730
+ def as_waveform_array(audio):
731
+ """Transpose one batch item of a pipeline's `.audios` output to (samples, channels).
732
+
733
+ `.audios` is shaped (batch, channels, samples). Under the diffusers default
734
+ `output_type='np'`, pipelines such as AudioLDM2 and StableAudio already call
735
+ `.numpy()` before returning, so each item here is a numpy ndarray rather than a
736
+ torch tensor - it has no `.float()`/`.cpu()` methods, only `.T`/`.astype()`.
737
+
738
+ Args:
739
+ audio: One batch item, shaped (channels, samples), as a torch tensor or numpy array
740
+
741
+ Returns:
742
+ Numpy float32 array shaped (samples, channels)
743
+ """
744
+ if isinstance(audio, torch.Tensor):
745
+ return audio.T.float().cpu().numpy()
746
+
747
+ return numpy.asarray(audio).T.astype(numpy.float32, copy=False)
748
+
749
+
750
+ def as_audio_track(audio):
751
+ """Convert a generated waveform into the tensor encode_video expects.
752
+
753
+ encode_video wants a float torch tensor on the CPU shaped (channels, samples) -
754
+ pipelines hand back bfloat16 tensors that are still on the GPU, or numpy arrays.
755
+
756
+ Args:
757
+ audio: Waveform as a torch tensor or numpy array
758
+
759
+ Returns:
760
+ Float CPU torch tensor holding the waveform
761
+ """
762
+ if not isinstance(audio, torch.Tensor):
763
+ audio = torch.from_numpy(numpy.asarray(audio))
764
+
765
+ return audio.detach().float().cpu()
766
+
767
+
768
+ def write_audio(output_path, waveform, sample_rate, **write_arguments):
769
+ """Write a single waveform to disk.
770
+
771
+ soundfile.write() hands the whole waveform to libsndfile in one call, whose vorbis
772
+ encoder segfaults past 2**21 frames - about 48 seconds of 44.1kHz audio. Writing in
773
+ chunks avoids that and bounds the encoder's working set for long audio.
774
+
775
+ Args:
776
+ output_path: Path of the file to write
777
+ waveform: Numpy array shaped (samples,) or (samples, channels)
778
+ sample_rate: Sample rate to record in the file
779
+ write_arguments: Container and encoding arguments for soundfile
780
+ """
781
+ channels = 1 if waveform.ndim == 1 else waveform.shape[1]
782
+
783
+ with soundfile.SoundFile(
784
+ output_path,
785
+ "w",
786
+ samplerate=sample_rate,
787
+ channels=channels,
788
+ **write_arguments,
789
+ ) as audio_file:
790
+ for start in range(0, len(waveform), AUDIO_WRITE_CHUNK_FRAMES):
791
+ audio_file.write(waveform[start : start + AUDIO_WRITE_CHUNK_FRAMES])
792
+
793
+
794
+ def normalize_audio(artifact):
795
+ """Convert an audio artifact into waveforms soundfile can write.
796
+
797
+ Pipelines return audio as torch tensors or numpy arrays, channels first and
798
+ optionally batched. soundfile wants samples first, one waveform at a time.
799
+
800
+ Args:
801
+ artifact: Audio waveform(s) as a torch tensor or numpy array
802
+
803
+ Returns:
804
+ List of numpy arrays shaped (samples,) or (samples, channels)
805
+ """
806
+ # Torch tensors may be on the GPU and in a dtype numpy does not understand
807
+ if hasattr(artifact, "detach"):
808
+ artifact = artifact.detach().float().cpu().numpy()
809
+
810
+ waveform = numpy.asarray(artifact)
811
+
812
+ if waveform.ndim == 1:
813
+ return [waveform]
814
+
815
+ if waveform.ndim == 2:
816
+ # Channels first - a waveform always has far more samples than channels
817
+ if waveform.shape[0] < waveform.shape[1]:
818
+ waveform = waveform.T
819
+ return [waveform]
820
+
821
+ if waveform.ndim == 3:
822
+ # (batch, channels, samples) - one waveform per batch item
823
+ return [item.T for item in waveform]
824
+
825
+ raise ValueError(f"Cannot save audio with shape {waveform.shape}")
826
+
827
+
828
+ def guess_extension(content_type):
829
+ """Determine file extension from MIME type.
830
+
831
+ Args:
832
+ content_type: MIME type string
833
+
834
+ Returns:
835
+ String containing file extension with leading dot
836
+ """
837
+ if not content_type:
838
+ logger.warning("No content type provided for extension guess")
839
+ return ""
840
+
841
+ # Audio is looked up first - soundfile picks the container from the extension and
842
+ # does not recognize every extension mimetypes suggests, such as '.oga' for ogg
843
+ if content_type in AUDIO_FORMATS:
844
+ return AUDIO_FORMATS[content_type][0]
845
+
846
+ ext = mimetypes.guess_extension(content_type)
847
+ if ext is not None:
848
+ return ext
849
+
850
+ return ""