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/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 ""
|