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
@@ -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
+ )