diffusers-workflow 0.4.0a3__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
- diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
- diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
- diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
- dw/__init__.py +353 -0
- dw/arguments.py +906 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +145 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/events.py +78 -0
- dw/hub_cache.py +289 -0
- dw/introspection.py +458 -0
- dw/log_setup.py +45 -0
- dw/pipeline_processors/chain.py +750 -0
- dw/pipeline_processors/config_objects.py +235 -0
- dw/pipeline_processors/pipeline.py +1687 -0
- dw/pipeline_processors/remote.py +18 -0
- dw/previous_results.py +259 -0
- dw/prompt_weighting.py +378 -0
- dw/repl.py +298 -0
- dw/repl_commands.py +808 -0
- dw/repl_worker.py +129 -0
- dw/result.py +850 -0
- dw/run.py +92 -0
- dw/schema.py +24 -0
- dw/security.py +379 -0
- dw/serve.py +70 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +588 -0
- dw/server/jobs.py +547 -0
- dw/server/ui/assets/abap-08VXUWAP.js +1 -0
- dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
- dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
- dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
- dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
- dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
- dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
- dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
- dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
- dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
- dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
- dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
- dw/server/ui/assets/css-DIMkf-bt.js +3 -0
- dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
- dw/server/ui/assets/cssMode-CEh6hWi2.js +1 -0
- dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
- dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
- dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
- dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
- dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
- dw/server/ui/assets/editor.api-CExg3_mM.js +847 -0
- dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
- dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
- dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
- dw/server/ui/assets/freemarker2-DH6orYh2.js +3 -0
- dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
- dw/server/ui/assets/go-C-y9NEjX.js +1 -0
- dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
- dw/server/ui/assets/handlebars-CbrMVW4Q.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-YDNPZw2M.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
- dw/server/ui/assets/index-B7-VcYS-.css +1 -0
- dw/server/ui/assets/index-D_EiPU3b.js +13 -0
- dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
- dw/server/ui/assets/java-BEtHBSE6.js +1 -0
- dw/server/ui/assets/javascript-dYuBvioq.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-CUqLM39V.js +7 -0
- dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
- dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
- dw/server/ui/assets/less-B9JPFI3C.js +2 -0
- dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
- dw/server/ui/assets/liquid-D6vxBzMv.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.js +4 -0
- dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
- dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
- dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
- dw/server/ui/assets/mdx-SHQb6vmD.js +1 -0
- dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
- dw/server/ui/assets/monaco--ixms01u.css +1 -0
- dw/server/ui/assets/monaco-CP-s5rcP.js +56 -0
- dw/server/ui/assets/msdax-DauUninz.js +1 -0
- dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
- dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
- dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
- dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
- dw/server/ui/assets/perl-oz_6vUea.js +1 -0
- dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
- dw/server/ui/assets/php-nr791fC2.js +1 -0
- dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
- dw/server/ui/assets/postiats-43DmfD33.js +1 -0
- dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
- dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
- dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
- dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
- dw/server/ui/assets/python-x0_EGHq9.js +1 -0
- dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
- dw/server/ui/assets/r-BwWrilGY.js +1 -0
- dw/server/ui/assets/razor-BZC4LQDP.js +1 -0
- dw/server/ui/assets/redis-ClamHrr6.js +1 -0
- dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
- dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
- dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
- dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
- dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
- dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
- dw/server/ui/assets/scheme-BeGwcela.js +1 -0
- dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
- dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
- dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
- dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
- dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
- dw/server/ui/assets/sql-NEE52Syq.js +1 -0
- dw/server/ui/assets/st-DbInun42.js +1 -0
- dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
- dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
- dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
- dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
- dw/server/ui/assets/tsMode-BTfA6SbD.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-CWA4MsNk.js +1 -0
- dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
- dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
- dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
- dw/server/ui/assets/workers-CWU0uvj5.js +1 -0
- dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
- dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/settings.py +77 -0
- dw/step.py +132 -0
- dw/tasks/audio_utils.py +266 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/concat_videos.py +80 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +139 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +661 -0
- dw/tasks/interpolate_frames.py +227 -0
- dw/tasks/model_cache.py +39 -0
- dw/tasks/pair_audio.py +58 -0
- dw/tasks/qr_code.py +19 -0
- dw/tasks/restore_faces.py +175 -0
- dw/tasks/rife_model.py +192 -0
- dw/tasks/segment.py +121 -0
- dw/tasks/task.py +474 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +168 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +154 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +376 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +68 -0
- dw/validate.py +43 -0
- dw/variables.py +153 -0
- dw/worker.py +517 -0
- dw/workflow.py +553 -0
- dw/workflow_schema.json +1157 -0
- dw/workflows/augment_prompt.json +65 -0
- dw/workflows/describe_image.json +58 -0
- dw/workflows/h3_context_ir.json +57 -0
- dw/workflows/test.json +31 -0
|
@@ -0,0 +1,750 @@
|
|
|
1
|
+
"""Segment-chained pipeline execution - videos of arbitrary length from
|
|
2
|
+
pipelines that generate short clips.
|
|
3
|
+
|
|
4
|
+
A "chain" block on a pipeline step runs the pipeline once per segment,
|
|
5
|
+
carries continuity from each segment into the next, trims the duplicated
|
|
6
|
+
boundary frames, and stitches the segments' frames and audio into one video.
|
|
7
|
+
|
|
8
|
+
Two ways to carry continuity:
|
|
9
|
+
- last_frame - the last frame becomes the next segment's keyframe, which is
|
|
10
|
+
what a keyframe-conditioned pipeline takes
|
|
11
|
+
- last_segment - the previous segment's frames and the soundtrack generated
|
|
12
|
+
with them become a video reference, which carries motion, camera and voice
|
|
13
|
+
across the seam rather than appearance alone
|
|
14
|
+
|
|
15
|
+
Two ways to specify the length:
|
|
16
|
+
- segments: N - run the pipeline N times as configured
|
|
17
|
+
- match_audio: true - derive the total frame count from the audio reference
|
|
18
|
+
in the step's arguments, slice that audio into frame-aligned per-segment
|
|
19
|
+
chunks, and mux the final video with the original, unsliced track - so the
|
|
20
|
+
soundtrack has no seams at all
|
|
21
|
+
|
|
22
|
+
The chain runs inside one cartesian iteration, so it composes with
|
|
23
|
+
previous_result fan-out: three keyframes in, three chained videos out.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
import gc
|
|
27
|
+
import logging
|
|
28
|
+
import math
|
|
29
|
+
import os
|
|
30
|
+
import sys
|
|
31
|
+
from dataclasses import dataclass
|
|
32
|
+
|
|
33
|
+
import numpy
|
|
34
|
+
import torch
|
|
35
|
+
from diffusers.utils import encode_video, is_av_available
|
|
36
|
+
|
|
37
|
+
from .. import empty_device_cache
|
|
38
|
+
from ..result import AudioVideo, get_artifact_list
|
|
39
|
+
from ..security import validate_output_path
|
|
40
|
+
from ..tasks.audio_utils import (
|
|
41
|
+
as_channels_samples,
|
|
42
|
+
equal_power_crossfade_join,
|
|
43
|
+
frames_to_samples,
|
|
44
|
+
slice_samples,
|
|
45
|
+
)
|
|
46
|
+
from ..tasks.video_utils import extract_frame, frames_as_pil_list
|
|
47
|
+
|
|
48
|
+
logger = logging.getLogger("dw")
|
|
49
|
+
|
|
50
|
+
# A runaway segment count is a configuration error - kept in the spirit of
|
|
51
|
+
# previous_results.MAX_ITERATIONS
|
|
52
|
+
MAX_SEGMENTS = 1000
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
class LastFrameContinuity:
|
|
56
|
+
"""Carry the last frame of each segment into the next as its keyframe."""
|
|
57
|
+
|
|
58
|
+
def __init__(self, config):
|
|
59
|
+
self.config = config
|
|
60
|
+
|
|
61
|
+
def extract(self, artifact):
|
|
62
|
+
return extract_frame(artifact, -1)
|
|
63
|
+
|
|
64
|
+
def inject(self, arguments, carry, segment_argument):
|
|
65
|
+
target = arguments.get(segment_argument)
|
|
66
|
+
if isinstance(target, list):
|
|
67
|
+
# A references list - the carry frame is appended as an image
|
|
68
|
+
# reference alongside the workflow's own references
|
|
69
|
+
arguments[segment_argument] = _with_carry_reference(target, carry)
|
|
70
|
+
else:
|
|
71
|
+
arguments[segment_argument] = carry
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class LastSegmentContinuity:
|
|
75
|
+
"""Carry the whole previous segment into the next as a video reference.
|
|
76
|
+
|
|
77
|
+
A still frame carries pose and colour and nothing else. A reference-
|
|
78
|
+
conditioned pipeline (MiniMax H3's ref2va) can take the previous segment
|
|
79
|
+
itself - its frames and the soundtrack generated with them - which carries
|
|
80
|
+
motion, camera and voice across the seam instead of just appearance.
|
|
81
|
+
|
|
82
|
+
Only the tail of the segment is worth carrying: 'carry_frames' bounds it,
|
|
83
|
+
and the soundtrack is cut to the same span so the reference's own audio and
|
|
84
|
+
video stay aligned. The generated media is already at the pipeline's own
|
|
85
|
+
rates, so the reference declares no rate of its own and nothing is
|
|
86
|
+
resampled on the way back in.
|
|
87
|
+
"""
|
|
88
|
+
|
|
89
|
+
def __init__(self, config):
|
|
90
|
+
self.config = config
|
|
91
|
+
|
|
92
|
+
def extract(self, artifact):
|
|
93
|
+
frames = frames_as_pil_list(artifact)
|
|
94
|
+
audio, sample_rate = _generated_audio(artifact)
|
|
95
|
+
|
|
96
|
+
carry_frames = self.config.carry_frames
|
|
97
|
+
if carry_frames is not None and carry_frames < len(frames):
|
|
98
|
+
if audio is not None:
|
|
99
|
+
if self.config.fps is None:
|
|
100
|
+
raise ValueError(
|
|
101
|
+
"Trimming a last_segment carry needs the frame rate - "
|
|
102
|
+
"set 'fps' on the chain or a 'frame_rate' pipeline argument"
|
|
103
|
+
)
|
|
104
|
+
samples = frames_to_samples(carry_frames, self.config.fps, sample_rate)
|
|
105
|
+
audio = audio[:, -samples:]
|
|
106
|
+
frames = frames[-carry_frames:]
|
|
107
|
+
|
|
108
|
+
if not self.config.carry_audio:
|
|
109
|
+
audio, sample_rate = None, None
|
|
110
|
+
|
|
111
|
+
return _SegmentCarry(frames, audio, sample_rate)
|
|
112
|
+
|
|
113
|
+
def inject(self, arguments, carry, segment_argument):
|
|
114
|
+
target = arguments.get(segment_argument)
|
|
115
|
+
if not isinstance(target, list):
|
|
116
|
+
raise ValueError(
|
|
117
|
+
f"The 'last_segment' continuity carries a video reference, so "
|
|
118
|
+
f"'{segment_argument}' must be a references list, not a "
|
|
119
|
+
f"{type(target).__name__}"
|
|
120
|
+
)
|
|
121
|
+
arguments[segment_argument] = _with_carry_video(target, carry)
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
CONTINUITY_MODES = {
|
|
125
|
+
"last_frame": LastFrameContinuity,
|
|
126
|
+
"last_segment": LastSegmentContinuity,
|
|
127
|
+
}
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
@dataclass
|
|
131
|
+
class _SegmentCarry:
|
|
132
|
+
"""The part of a finished segment a last_segment chain conditions on."""
|
|
133
|
+
|
|
134
|
+
frames: list
|
|
135
|
+
audio: object # (channels, samples) numpy, or None
|
|
136
|
+
sample_rate: int
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
@dataclass
|
|
140
|
+
class Segment:
|
|
141
|
+
"""One planned pipeline invocation of a chain."""
|
|
142
|
+
|
|
143
|
+
index: int
|
|
144
|
+
num_frames: int # frames this segment generates; None - use the step's own
|
|
145
|
+
audio_start_frame: int # generation-timeline frame its audio slice starts at
|
|
146
|
+
head_trim: int # frames dropped from its head on the output timeline
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
class SegmentedFrames:
|
|
150
|
+
"""Chained segments spilled to disk, replayed at save time.
|
|
151
|
+
|
|
152
|
+
Iterating yields one uint8 (frames, height, width, 3) torch tensor per
|
|
153
|
+
segment file - the chunk shape encode_video streams from - so the final
|
|
154
|
+
video is written holding only one segment in memory at a time.
|
|
155
|
+
"""
|
|
156
|
+
|
|
157
|
+
def __init__(self, paths, total_frames=None, keep_files=False):
|
|
158
|
+
"""
|
|
159
|
+
Args:
|
|
160
|
+
paths: The segment files, in output order
|
|
161
|
+
total_frames: Frames to yield in total - the match_audio tail trim.
|
|
162
|
+
None yields every stored frame
|
|
163
|
+
keep_files: Leave the segment files in place after cleanup()
|
|
164
|
+
"""
|
|
165
|
+
self.paths = list(paths)
|
|
166
|
+
self.total_frames = total_frames
|
|
167
|
+
self.keep_files = keep_files
|
|
168
|
+
|
|
169
|
+
def __len__(self):
|
|
170
|
+
"""Chunk count - one per segment file."""
|
|
171
|
+
return len(self.paths)
|
|
172
|
+
|
|
173
|
+
def __iter__(self):
|
|
174
|
+
remaining = self.total_frames
|
|
175
|
+
for path in self.paths:
|
|
176
|
+
frames = _decode_segment(path)
|
|
177
|
+
if remaining is not None:
|
|
178
|
+
frames = frames[:remaining]
|
|
179
|
+
remaining -= len(frames)
|
|
180
|
+
if len(frames):
|
|
181
|
+
yield frames
|
|
182
|
+
if remaining == 0:
|
|
183
|
+
return
|
|
184
|
+
|
|
185
|
+
def cleanup(self):
|
|
186
|
+
"""Remove the segment files once the final video is safely written."""
|
|
187
|
+
if self.keep_files:
|
|
188
|
+
return
|
|
189
|
+
for path in self.paths:
|
|
190
|
+
try:
|
|
191
|
+
os.remove(path)
|
|
192
|
+
except FileNotFoundError:
|
|
193
|
+
pass
|
|
194
|
+
|
|
195
|
+
|
|
196
|
+
class SegmentSpill:
|
|
197
|
+
"""Writes each completed segment to disk as a playable mp4.
|
|
198
|
+
|
|
199
|
+
Files are named {prefix}.{iteration}.segment-{index:03d}.mp4 in the
|
|
200
|
+
workflow's output directory - a crashed chain leaves them behind, ready to
|
|
201
|
+
salvage with gather_videos + concat_videos.
|
|
202
|
+
"""
|
|
203
|
+
|
|
204
|
+
def __init__(self, pipeline, config):
|
|
205
|
+
output_dir = getattr(pipeline, "output_dir", None)
|
|
206
|
+
file_prefix = getattr(pipeline, "file_prefix", None)
|
|
207
|
+
if not output_dir or not file_prefix:
|
|
208
|
+
raise ValueError(
|
|
209
|
+
"save_segments needs the workflow's output directory - it is "
|
|
210
|
+
"only available when the chain runs through a workflow"
|
|
211
|
+
)
|
|
212
|
+
if not is_av_available():
|
|
213
|
+
raise ValueError(
|
|
214
|
+
"save_segments writes mp4 segment files with PyAV - install "
|
|
215
|
+
"it with: pip install av"
|
|
216
|
+
)
|
|
217
|
+
if config.fps is None:
|
|
218
|
+
raise ValueError(
|
|
219
|
+
"save_segments needs the frame rate to encode segment files - "
|
|
220
|
+
"set 'fps' on the chain or a 'frame_rate' pipeline argument"
|
|
221
|
+
)
|
|
222
|
+
|
|
223
|
+
# The same step can chain more than once (previous_result fan-out) -
|
|
224
|
+
# a per-wrapper counter keeps each iteration's files apart
|
|
225
|
+
iteration = getattr(pipeline, "_chain_iteration", -1) + 1
|
|
226
|
+
pipeline._chain_iteration = iteration
|
|
227
|
+
|
|
228
|
+
self.output_dir = output_dir
|
|
229
|
+
self.base_name = f"{file_prefix}.{iteration}"
|
|
230
|
+
self.fps = config.fps
|
|
231
|
+
self.paths = []
|
|
232
|
+
|
|
233
|
+
def write(self, frames, audio, sample_rate):
|
|
234
|
+
"""Encode one trimmed segment to disk and record its path.
|
|
235
|
+
|
|
236
|
+
Args:
|
|
237
|
+
frames: The segment's on-timeline PIL frames
|
|
238
|
+
audio: The segment's on-timeline generated audio as
|
|
239
|
+
(channels, samples) numpy, or None - muxed in so a crashed
|
|
240
|
+
chain leaves fully playable segments
|
|
241
|
+
sample_rate: Sample rate of that audio
|
|
242
|
+
"""
|
|
243
|
+
path = validate_output_path(
|
|
244
|
+
os.path.join(
|
|
245
|
+
self.output_dir,
|
|
246
|
+
f"{self.base_name}.segment-{len(self.paths):03d}.mp4",
|
|
247
|
+
),
|
|
248
|
+
self.output_dir,
|
|
249
|
+
)
|
|
250
|
+
|
|
251
|
+
audio_track = None
|
|
252
|
+
if audio is not None and audio.shape[1] and sample_rate is not None:
|
|
253
|
+
audio_track = torch.from_numpy(numpy.ascontiguousarray(audio))
|
|
254
|
+
|
|
255
|
+
encode_video(
|
|
256
|
+
frames,
|
|
257
|
+
fps=self.fps,
|
|
258
|
+
output_path=path,
|
|
259
|
+
audio=audio_track,
|
|
260
|
+
audio_sample_rate=sample_rate if audio_track is not None else None,
|
|
261
|
+
)
|
|
262
|
+
self.paths.append(path)
|
|
263
|
+
logger.info(f"Saved chain segment to {path}")
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def _decode_segment(path):
|
|
267
|
+
"""Read a segment file back as a uint8 (frames, height, width, 3) tensor."""
|
|
268
|
+
import av
|
|
269
|
+
|
|
270
|
+
with av.open(path) as container:
|
|
271
|
+
frames = [
|
|
272
|
+
frame.to_ndarray(format="rgb24") for frame in container.decode(video=0)
|
|
273
|
+
]
|
|
274
|
+
return torch.from_numpy(numpy.stack(frames, axis=0))
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
def run_chain(pipeline, chain_definition, arguments):
|
|
278
|
+
"""Run a pipeline's chain and stitch the segments into one video.
|
|
279
|
+
|
|
280
|
+
Args:
|
|
281
|
+
pipeline: The loaded Pipeline wrapper - each segment goes through its
|
|
282
|
+
_run_once, so prompt handling matches an unchained run
|
|
283
|
+
chain_definition: The step's "chain" block
|
|
284
|
+
arguments: Fully resolved arguments for one iteration of the step
|
|
285
|
+
|
|
286
|
+
Returns:
|
|
287
|
+
A single AudioVideo holding the stitched frames, and either the joined
|
|
288
|
+
generated audio, the original match_audio track, or no audio at all
|
|
289
|
+
"""
|
|
290
|
+
# Step.run resolved any previous_result references in the chain's prompts,
|
|
291
|
+
# which the chain block cannot express on its own
|
|
292
|
+
config = ChainConfig(
|
|
293
|
+
chain_definition, arguments, getattr(pipeline, "chain_prompts", None)
|
|
294
|
+
)
|
|
295
|
+
continuity = CONTINUITY_MODES[config.continuity](config)
|
|
296
|
+
|
|
297
|
+
# With save_segments, each completed segment is written to disk and its
|
|
298
|
+
# frames freed, bounding memory to one segment - a crash leaves the
|
|
299
|
+
# finished segments behind as playable files
|
|
300
|
+
spill = SegmentSpill(pipeline, config) if config.save_segments else None
|
|
301
|
+
|
|
302
|
+
frames = [] # PIL frames on the output timeline (unspilled chains)
|
|
303
|
+
audio = None # joined generated audio, (channels, samples) float32
|
|
304
|
+
audio_rate = None
|
|
305
|
+
carry = None
|
|
306
|
+
|
|
307
|
+
for segment in config.plan:
|
|
308
|
+
segment_arguments = dict(arguments)
|
|
309
|
+
|
|
310
|
+
if config.prompts:
|
|
311
|
+
segment_arguments["prompt"] = config.prompts[
|
|
312
|
+
min(segment.index, len(config.prompts) - 1)
|
|
313
|
+
]
|
|
314
|
+
|
|
315
|
+
if config.source_audio is not None:
|
|
316
|
+
segment_arguments["num_frames"] = segment.num_frames
|
|
317
|
+
segment_arguments["references"] = _sliced_references(
|
|
318
|
+
config, segment, arguments["references"]
|
|
319
|
+
)
|
|
320
|
+
|
|
321
|
+
if segment.index > 0:
|
|
322
|
+
continuity.inject(segment_arguments, carry, config.segment_argument)
|
|
323
|
+
|
|
324
|
+
logger.info(
|
|
325
|
+
f"Chain segment {segment.index + 1}/{len(config.plan)}"
|
|
326
|
+
+ (f": {segment.num_frames} frames" if segment.num_frames else "")
|
|
327
|
+
)
|
|
328
|
+
|
|
329
|
+
output = pipeline._run_once(segment_arguments)
|
|
330
|
+
artifact = _single_artifact(output)
|
|
331
|
+
|
|
332
|
+
carry = continuity.extract(artifact)
|
|
333
|
+
segment_frames = frames_as_pil_list(artifact)
|
|
334
|
+
segment_audio, segment_rate = _generated_audio(artifact)
|
|
335
|
+
|
|
336
|
+
kept_frames = segment_frames[segment.head_trim :]
|
|
337
|
+
if spill is not None:
|
|
338
|
+
spill.write(
|
|
339
|
+
kept_frames,
|
|
340
|
+
_on_timeline_audio(segment_audio, segment, config, segment_rate),
|
|
341
|
+
segment_rate,
|
|
342
|
+
)
|
|
343
|
+
else:
|
|
344
|
+
frames.extend(kept_frames)
|
|
345
|
+
|
|
346
|
+
if config.source_audio is None and segment_audio is not None:
|
|
347
|
+
audio, audio_rate = _joined_audio(
|
|
348
|
+
audio, audio_rate, segment_audio, segment_rate, segment, config
|
|
349
|
+
)
|
|
350
|
+
|
|
351
|
+
# The segment's raw output is finished with - the frames live on
|
|
352
|
+
# (in RAM or on disk) and the carry frame is extracted. Free it
|
|
353
|
+
# before the next segment needs the accelerator.
|
|
354
|
+
del output, artifact, segment_frames, segment_audio, kept_frames
|
|
355
|
+
gc.collect()
|
|
356
|
+
empty_device_cache()
|
|
357
|
+
|
|
358
|
+
if spill is not None:
|
|
359
|
+
# match_audio overshoots by design - the tail trim happens as the
|
|
360
|
+
# lazy frames replay, so the files themselves stay whole
|
|
361
|
+
frames = SegmentedFrames(
|
|
362
|
+
spill.paths,
|
|
363
|
+
config.total_frames if config.source_audio is not None else None,
|
|
364
|
+
config.keep_segments,
|
|
365
|
+
)
|
|
366
|
+
|
|
367
|
+
if config.source_audio is not None:
|
|
368
|
+
# The video matches the track's duration; the original, unsliced audio
|
|
369
|
+
# is muxed in so the soundtrack has no seams
|
|
370
|
+
if spill is None:
|
|
371
|
+
frames = frames[: config.total_frames]
|
|
372
|
+
return AudioVideo(frames, config.source_audio, config.source_rate)
|
|
373
|
+
|
|
374
|
+
return AudioVideo(frames, audio, audio_rate)
|
|
375
|
+
|
|
376
|
+
|
|
377
|
+
class ChainConfig:
|
|
378
|
+
"""Validated chain settings plus the planned segments for one run."""
|
|
379
|
+
|
|
380
|
+
def __init__(self, chain_definition, arguments, resolved_prompts=None):
|
|
381
|
+
segments = chain_definition.get("segments", None)
|
|
382
|
+
match_audio = bool(chain_definition.get("match_audio", False))
|
|
383
|
+
if (segments is not None) == match_audio:
|
|
384
|
+
raise ValueError("A chain needs exactly one of 'segments' or 'match_audio'")
|
|
385
|
+
|
|
386
|
+
self.continuity = chain_definition.get("continuity", "last_frame")
|
|
387
|
+
if self.continuity not in CONTINUITY_MODES:
|
|
388
|
+
known = ", ".join(sorted(CONTINUITY_MODES))
|
|
389
|
+
raise ValueError(
|
|
390
|
+
f"Unknown chain continuity '{self.continuity}' - expected one of {known}"
|
|
391
|
+
)
|
|
392
|
+
|
|
393
|
+
self.segment_argument = chain_definition.get("segment_argument", "image")
|
|
394
|
+
self.carry_frames = chain_definition.get("carry_frames", None)
|
|
395
|
+
if self.carry_frames is not None:
|
|
396
|
+
self.carry_frames = int(self.carry_frames)
|
|
397
|
+
if self.carry_frames < 1:
|
|
398
|
+
raise ValueError(
|
|
399
|
+
f"Chain 'carry_frames' must be at least 1, got {self.carry_frames}"
|
|
400
|
+
)
|
|
401
|
+
self.carry_audio = bool(chain_definition.get("carry_audio", True))
|
|
402
|
+
self.trim_frames = int(chain_definition.get("trim_frames", 1))
|
|
403
|
+
self.crossfade_ms = float(chain_definition.get("crossfade_ms", 75))
|
|
404
|
+
self.prompts = resolved_prompts or chain_definition.get("prompts", None)
|
|
405
|
+
self.fps = _resolve_fps(chain_definition, arguments)
|
|
406
|
+
self.frame_snap = chain_definition.get("frame_snap", None)
|
|
407
|
+
self.save_segments = bool(chain_definition.get("save_segments", False))
|
|
408
|
+
self.keep_segments = bool(chain_definition.get("keep_segments", False))
|
|
409
|
+
|
|
410
|
+
self.source_audio = None
|
|
411
|
+
self.source_rate = None
|
|
412
|
+
self.audio_reference = None
|
|
413
|
+
self.total_frames = None
|
|
414
|
+
|
|
415
|
+
if match_audio:
|
|
416
|
+
self._plan_from_audio(arguments)
|
|
417
|
+
else:
|
|
418
|
+
segments = int(segments)
|
|
419
|
+
if not 1 <= segments <= MAX_SEGMENTS:
|
|
420
|
+
raise ValueError(
|
|
421
|
+
f"Chain 'segments' must be between 1 and {MAX_SEGMENTS}, got {segments}"
|
|
422
|
+
)
|
|
423
|
+
num_frames = arguments.get("num_frames", None)
|
|
424
|
+
if num_frames is not None:
|
|
425
|
+
validate_frame_snap(int(num_frames), self.frame_snap)
|
|
426
|
+
self.plan = [
|
|
427
|
+
Segment(
|
|
428
|
+
index,
|
|
429
|
+
int(num_frames) if num_frames is not None else None,
|
|
430
|
+
0,
|
|
431
|
+
self.trim_frames if index > 0 else 0,
|
|
432
|
+
)
|
|
433
|
+
for index in range(segments)
|
|
434
|
+
]
|
|
435
|
+
|
|
436
|
+
def _plan_from_audio(self, arguments):
|
|
437
|
+
"""Derive the segment plan from the audio reference's duration."""
|
|
438
|
+
if self.fps is None:
|
|
439
|
+
raise ValueError(
|
|
440
|
+
"A match_audio chain needs the frame rate - set 'fps' on the "
|
|
441
|
+
"chain or a 'frame_rate' pipeline argument"
|
|
442
|
+
)
|
|
443
|
+
|
|
444
|
+
num_frames = arguments.get("num_frames", None)
|
|
445
|
+
if num_frames is None:
|
|
446
|
+
raise ValueError(
|
|
447
|
+
"A match_audio chain needs 'num_frames' in the step's arguments "
|
|
448
|
+
"as the per-segment length"
|
|
449
|
+
)
|
|
450
|
+
|
|
451
|
+
reference = _find_audio_reference(arguments)
|
|
452
|
+
self.audio_reference = reference
|
|
453
|
+
self.source_audio = as_channels_samples(reference.audio)
|
|
454
|
+
self.source_rate = reference.sample_rate
|
|
455
|
+
if self.source_rate is None:
|
|
456
|
+
raise ValueError("The chain's audio reference has no sample rate")
|
|
457
|
+
|
|
458
|
+
total_samples = self.source_audio.shape[1]
|
|
459
|
+
self.total_frames = max(1, round(total_samples / self.source_rate * self.fps))
|
|
460
|
+
self.plan = plan_segments(
|
|
461
|
+
self.total_frames, int(num_frames), self.trim_frames, self.frame_snap
|
|
462
|
+
)
|
|
463
|
+
duration = total_samples / self.source_rate
|
|
464
|
+
logger.info(
|
|
465
|
+
f"Chaining to match {duration:.2f}s of audio: {self.total_frames} "
|
|
466
|
+
f"frames across {len(self.plan)} segments"
|
|
467
|
+
)
|
|
468
|
+
|
|
469
|
+
|
|
470
|
+
def plan_segments(total_frames, segment_frames, trim_frames, frame_snap=None):
|
|
471
|
+
"""Plan the segments that cover a total frame count.
|
|
472
|
+
|
|
473
|
+
Every segment generates segment_frames frames except possibly the last,
|
|
474
|
+
which shrinks to what remains - snapped up to a count the pipeline accepts.
|
|
475
|
+
Each segment after the first has trim_frames dropped from its head, so it
|
|
476
|
+
contributes segment_frames - trim_frames new frames to the output.
|
|
477
|
+
|
|
478
|
+
Args:
|
|
479
|
+
total_frames: Frames the stitched output must cover
|
|
480
|
+
segment_frames: Frames a full segment generates
|
|
481
|
+
trim_frames: Head frames dropped from every segment after the first
|
|
482
|
+
frame_snap: Optional dict with modulus/remainder and min/max_frames
|
|
483
|
+
describing the counts the pipeline accepts
|
|
484
|
+
|
|
485
|
+
Returns:
|
|
486
|
+
List of Segment
|
|
487
|
+
"""
|
|
488
|
+
if segment_frames <= trim_frames:
|
|
489
|
+
raise ValueError(
|
|
490
|
+
f"Segments of {segment_frames} frames cannot progress past a head "
|
|
491
|
+
f"trim of {trim_frames} frames"
|
|
492
|
+
)
|
|
493
|
+
validate_frame_snap(segment_frames, frame_snap)
|
|
494
|
+
|
|
495
|
+
plan = []
|
|
496
|
+
covered = 0
|
|
497
|
+
while covered < total_frames:
|
|
498
|
+
if len(plan) >= MAX_SEGMENTS:
|
|
499
|
+
raise ValueError(f"Chain would exceed {MAX_SEGMENTS} segments")
|
|
500
|
+
|
|
501
|
+
head_trim = trim_frames if plan else 0
|
|
502
|
+
needed = (total_frames - covered) + head_trim
|
|
503
|
+
if needed >= segment_frames:
|
|
504
|
+
num_frames = segment_frames
|
|
505
|
+
else:
|
|
506
|
+
# The last segment generates only what remains, snapped up to a
|
|
507
|
+
# count the pipeline accepts; the overshoot is trimmed at the end
|
|
508
|
+
num_frames = snap_frames(needed, frame_snap)
|
|
509
|
+
|
|
510
|
+
plan.append(Segment(len(plan), num_frames, covered - head_trim, head_trim))
|
|
511
|
+
covered += num_frames - head_trim
|
|
512
|
+
|
|
513
|
+
return plan
|
|
514
|
+
|
|
515
|
+
|
|
516
|
+
def snap_frames(count, frame_snap):
|
|
517
|
+
"""The smallest frame count the pipeline accepts that covers count."""
|
|
518
|
+
if not frame_snap:
|
|
519
|
+
return count
|
|
520
|
+
|
|
521
|
+
modulus = frame_snap["modulus"]
|
|
522
|
+
remainder = frame_snap["remainder"]
|
|
523
|
+
target = max(count, frame_snap.get("min_frames", 1))
|
|
524
|
+
|
|
525
|
+
steps = max(0, math.ceil((target - remainder) / modulus))
|
|
526
|
+
snapped = steps * modulus + remainder
|
|
527
|
+
while snapped < target:
|
|
528
|
+
snapped += modulus
|
|
529
|
+
|
|
530
|
+
max_frames = frame_snap.get("max_frames", None)
|
|
531
|
+
if max_frames is not None and snapped > max_frames:
|
|
532
|
+
raise ValueError(
|
|
533
|
+
f"Cannot snap {count} frames into the pipeline's accepted range - "
|
|
534
|
+
f"the next valid count {snapped} exceeds max_frames {max_frames}"
|
|
535
|
+
)
|
|
536
|
+
return snapped
|
|
537
|
+
|
|
538
|
+
|
|
539
|
+
def validate_frame_snap(num_frames, frame_snap):
|
|
540
|
+
"""Check a configured num_frames against the pipeline's constraint."""
|
|
541
|
+
if not frame_snap:
|
|
542
|
+
return
|
|
543
|
+
|
|
544
|
+
modulus = frame_snap["modulus"]
|
|
545
|
+
remainder = frame_snap["remainder"]
|
|
546
|
+
problems = []
|
|
547
|
+
if (num_frames - remainder) % modulus != 0:
|
|
548
|
+
problems.append(f"counts must be {modulus}*n+{remainder}")
|
|
549
|
+
min_frames = frame_snap.get("min_frames", None)
|
|
550
|
+
if min_frames is not None and num_frames < min_frames:
|
|
551
|
+
problems.append(f"at least {min_frames}")
|
|
552
|
+
max_frames = frame_snap.get("max_frames", None)
|
|
553
|
+
if max_frames is not None and num_frames > max_frames:
|
|
554
|
+
problems.append(f"at most {max_frames}")
|
|
555
|
+
|
|
556
|
+
if problems:
|
|
557
|
+
raise ValueError(
|
|
558
|
+
f"num_frames {num_frames} does not satisfy the pipeline's frame "
|
|
559
|
+
f"constraint: {'; '.join(problems)}"
|
|
560
|
+
)
|
|
561
|
+
|
|
562
|
+
|
|
563
|
+
def _resolve_fps(chain_definition, arguments):
|
|
564
|
+
"""The frame rate used for audio math - explicit, or the pipeline's own."""
|
|
565
|
+
fps = chain_definition.get("fps", arguments.get("frame_rate", None))
|
|
566
|
+
return float(fps) if fps is not None else None
|
|
567
|
+
|
|
568
|
+
|
|
569
|
+
def _find_audio_reference(arguments):
|
|
570
|
+
"""The single audio reference a match_audio chain slices per segment."""
|
|
571
|
+
references = arguments.get("references", None)
|
|
572
|
+
if not isinstance(references, list):
|
|
573
|
+
raise ValueError(
|
|
574
|
+
"A match_audio chain needs a 'references' argument holding the "
|
|
575
|
+
"audio reference to match"
|
|
576
|
+
)
|
|
577
|
+
|
|
578
|
+
audio_references = [
|
|
579
|
+
reference
|
|
580
|
+
for reference in references
|
|
581
|
+
if getattr(reference, "kind", None) == "audio"
|
|
582
|
+
]
|
|
583
|
+
if len(audio_references) != 1:
|
|
584
|
+
raise ValueError(
|
|
585
|
+
f"A match_audio chain needs exactly one audio reference, "
|
|
586
|
+
f"found {len(audio_references)}"
|
|
587
|
+
)
|
|
588
|
+
return audio_references[0]
|
|
589
|
+
|
|
590
|
+
|
|
591
|
+
def _sliced_references(config, segment, references):
|
|
592
|
+
"""A copy of the references list with the segment's audio slice swapped in.
|
|
593
|
+
|
|
594
|
+
The original list and reference objects are never touched - iteration
|
|
595
|
+
arguments share nested values, so they must not be mutated in place.
|
|
596
|
+
"""
|
|
597
|
+
start = frames_to_samples(segment.audio_start_frame, config.fps, config.source_rate)
|
|
598
|
+
length = frames_to_samples(segment.num_frames, config.fps, config.source_rate)
|
|
599
|
+
piece = slice_samples(config.source_audio, start, length)
|
|
600
|
+
|
|
601
|
+
sliced = type(config.audio_reference)(
|
|
602
|
+
audio=torch.from_numpy(piece), sample_rate=config.source_rate
|
|
603
|
+
)
|
|
604
|
+
return [
|
|
605
|
+
sliced if reference is config.audio_reference else reference
|
|
606
|
+
for reference in references
|
|
607
|
+
]
|
|
608
|
+
|
|
609
|
+
|
|
610
|
+
def _with_carry_reference(references, carry):
|
|
611
|
+
"""A copy of a references list with the carry frame appended as an image
|
|
612
|
+
reference of the same type the workflow already uses."""
|
|
613
|
+
image_reference = next(
|
|
614
|
+
(
|
|
615
|
+
reference
|
|
616
|
+
for reference in references
|
|
617
|
+
if getattr(reference, "kind", None) == "image"
|
|
618
|
+
),
|
|
619
|
+
None,
|
|
620
|
+
)
|
|
621
|
+
if image_reference is None:
|
|
622
|
+
raise ValueError(
|
|
623
|
+
"Cannot carry a frame into a references list that has no image "
|
|
624
|
+
"reference to model the new one on"
|
|
625
|
+
)
|
|
626
|
+
return list(references) + [type(image_reference)(image=carry)]
|
|
627
|
+
|
|
628
|
+
|
|
629
|
+
def _with_carry_video(references, carry):
|
|
630
|
+
"""A copy of a references list with the carry segment appended as a video
|
|
631
|
+
reference of the same family the workflow already uses."""
|
|
632
|
+
reference_type = _video_reference_type(references)
|
|
633
|
+
arguments = {"frames": carry.frames}
|
|
634
|
+
if carry.audio is not None:
|
|
635
|
+
arguments["audio"] = torch.from_numpy(carry.audio)
|
|
636
|
+
arguments["sample_rate"] = carry.sample_rate
|
|
637
|
+
return list(references) + [reference_type(**arguments)]
|
|
638
|
+
|
|
639
|
+
|
|
640
|
+
def _video_reference_type(references):
|
|
641
|
+
"""The video reference class of the family the workflow's references come from.
|
|
642
|
+
|
|
643
|
+
A workflow that already passes a video reference names the class outright.
|
|
644
|
+
Otherwise it is the video-kind class living beside the references it does
|
|
645
|
+
pass - the chain never imports a pipeline's reference types itself, the way
|
|
646
|
+
_with_carry_reference models its carry on the list it was given.
|
|
647
|
+
"""
|
|
648
|
+
if not references:
|
|
649
|
+
raise ValueError(
|
|
650
|
+
"Cannot carry a segment into an empty references list - a "
|
|
651
|
+
"last_segment chain needs the workflow's own references to model "
|
|
652
|
+
"the carry on"
|
|
653
|
+
)
|
|
654
|
+
|
|
655
|
+
for reference in references:
|
|
656
|
+
if getattr(reference, "kind", None) == "video":
|
|
657
|
+
return type(reference)
|
|
658
|
+
|
|
659
|
+
module = sys.modules.get(type(references[0]).__module__, None)
|
|
660
|
+
for candidate in vars(module).values() if module else ():
|
|
661
|
+
if isinstance(candidate, type) and getattr(candidate, "kind", None) == "video":
|
|
662
|
+
return candidate
|
|
663
|
+
|
|
664
|
+
raise ValueError(
|
|
665
|
+
f"Cannot carry a segment as a video reference - no video reference type "
|
|
666
|
+
f"found alongside {type(references[0]).__name__}"
|
|
667
|
+
)
|
|
668
|
+
|
|
669
|
+
|
|
670
|
+
def _single_artifact(output):
|
|
671
|
+
"""The one video artifact a chain segment must produce.
|
|
672
|
+
|
|
673
|
+
Modular pipelines asked for extra outputs return them alongside the video -
|
|
674
|
+
those are dropped here. More than one video means batched generation, which
|
|
675
|
+
a chain cannot stitch.
|
|
676
|
+
"""
|
|
677
|
+
artifacts = get_artifact_list(output)
|
|
678
|
+
videos = [artifact for artifact in artifacts if _is_video_artifact(artifact)]
|
|
679
|
+
|
|
680
|
+
if len(videos) != 1:
|
|
681
|
+
raise ValueError(
|
|
682
|
+
f"A chained pipeline must generate exactly one video per segment, "
|
|
683
|
+
f"got {len(videos)} - batched generation cannot be chained"
|
|
684
|
+
)
|
|
685
|
+
if len(artifacts) > 1:
|
|
686
|
+
logger.debug(f"Chain segment dropped {len(artifacts) - 1} non-video output(s)")
|
|
687
|
+
return videos[0]
|
|
688
|
+
|
|
689
|
+
|
|
690
|
+
def _is_video_artifact(artifact):
|
|
691
|
+
if isinstance(artifact, AudioVideo):
|
|
692
|
+
return True
|
|
693
|
+
if isinstance(artifact, list) and artifact:
|
|
694
|
+
return not isinstance(artifact[0], str)
|
|
695
|
+
return hasattr(artifact, "ndim") and artifact.ndim >= 3
|
|
696
|
+
|
|
697
|
+
|
|
698
|
+
def _generated_audio(artifact):
|
|
699
|
+
"""The audio generated with a segment, as (channels, samples) numpy."""
|
|
700
|
+
if isinstance(artifact, AudioVideo) and artifact.audio is not None:
|
|
701
|
+
return as_channels_samples(artifact.audio), artifact.sample_rate
|
|
702
|
+
return None, None
|
|
703
|
+
|
|
704
|
+
|
|
705
|
+
def _on_timeline_audio(segment_audio, segment, config, segment_rate):
|
|
706
|
+
"""The part of a segment's generated audio that survives the head trim.
|
|
707
|
+
|
|
708
|
+
Muxed into the segment's spill file so a crashed chain leaves fully
|
|
709
|
+
playable segments; the final soundtrack still comes from the accumulated
|
|
710
|
+
crossfaded track (or the original match_audio track).
|
|
711
|
+
"""
|
|
712
|
+
if segment_audio is None:
|
|
713
|
+
return None
|
|
714
|
+
trim_samples = frames_to_samples(segment.head_trim, config.fps, segment_rate)
|
|
715
|
+
return segment_audio[:, trim_samples:]
|
|
716
|
+
|
|
717
|
+
|
|
718
|
+
def _joined_audio(audio, audio_rate, segment_audio, segment_rate, segment, config):
|
|
719
|
+
"""Fold one segment's generated audio into the accumulated track.
|
|
720
|
+
|
|
721
|
+
The samples matching the segment's trimmed head frames are cut off and
|
|
722
|
+
used as crossfade material against the tail of the accumulated audio, so
|
|
723
|
+
the audio timeline shortens by exactly as much as the video's.
|
|
724
|
+
"""
|
|
725
|
+
if audio is None:
|
|
726
|
+
return segment_audio, segment_rate
|
|
727
|
+
|
|
728
|
+
if segment_rate != audio_rate:
|
|
729
|
+
raise ValueError(
|
|
730
|
+
f"Chain segments generated audio at different sample rates: "
|
|
731
|
+
f"{audio_rate} then {segment_rate}"
|
|
732
|
+
)
|
|
733
|
+
|
|
734
|
+
if segment.head_trim > 0 and config.fps is None:
|
|
735
|
+
raise ValueError(
|
|
736
|
+
"Joining generated audio needs the frame rate - set 'fps' on the "
|
|
737
|
+
"chain or a 'frame_rate' pipeline argument"
|
|
738
|
+
)
|
|
739
|
+
|
|
740
|
+
trim_samples = (
|
|
741
|
+
frames_to_samples(segment.head_trim, config.fps, audio_rate)
|
|
742
|
+
if segment.head_trim
|
|
743
|
+
else 0
|
|
744
|
+
)
|
|
745
|
+
head = segment_audio[:, :trim_samples]
|
|
746
|
+
body = segment_audio[:, trim_samples:]
|
|
747
|
+
return (
|
|
748
|
+
equal_power_crossfade_join(audio, head, body, audio_rate, config.crossfade_ms),
|
|
749
|
+
audio_rate,
|
|
750
|
+
)
|