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