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,228 @@
1
+ """
2
+ Speech generation via HuggingFace transformers.
3
+
4
+ Takes a line of text and speaks it with a local text-to-speech model, returning
5
+ the waveform and the rate it was generated at so the rest of the audio plumbing -
6
+ slice_audio, fade_audio and pair_audio - composes with it directly (concat_videos
7
+ and dissolve_videos join videos; pair the track onto a video first).
8
+
9
+ The role this is built for is voice *timbre reference*, not the track a mouth
10
+ follows. MiniMax H3 lip-syncs well when it generates the speech itself and poorly
11
+ when it must follow supplied audio, so its audio reference takes a few seconds of
12
+ a voice to fix timbre, pitch and delivery while the model still generates the
13
+ line. A generated clip referenced in every shot makes voice consistency an actual
14
+ conditioning signal rather than a prose description that has to land identically
15
+ a dozen times. The other honest uses are a voice that must be matched, and
16
+ narration over shots where nothing has to lip-sync to it.
17
+ """
18
+
19
+ import logging
20
+
21
+ import torch
22
+ from transformers import pipeline as hf_pipeline
23
+
24
+ from .. import preferred_task_dtype
25
+ from ..result import AudioTrack
26
+ from .audio_utils import as_channels_samples, load_audio, resample_waveform
27
+ from .model_cache import cached_model, hf_pipeline_placement
28
+
29
+ logger = logging.getLogger("dw")
30
+
31
+ # Small, loads without a separate speaker-embedding dataset, and its voice presets
32
+ # give distinct speakers - which is the point when two characters have to sound
33
+ # like two people. A single-voice model like facebook/mms-tts-eng (which takes no
34
+ # voice_preset) is a quarter the size and a reasonable override where only one
35
+ # voice is needed
36
+ _DEFAULT_MODEL = "suno/bark-small"
37
+
38
+ # What SpeechT5's voice-cloning recipes are built around - trained on the same
39
+ # corpus (VoxCeleb) the CMU ARCTIC x-vectors that ship with the model card come
40
+ # from, so a reference clip reduces to a vector in the space the model expects
41
+ _SPEAKER_ENCODER_MODEL = "speechbrain/spkrec-xvect-voxceleb"
42
+ # What the encoder was trained on - resampling anything else to this rate is
43
+ # part of extracting a voiceprint, not an approximation of one
44
+ _SPEAKER_ENCODER_SAMPLE_RATE = 16000
45
+
46
+
47
+ def _speaker_embedding_tensor(location, device, dtype):
48
+ """The x-vector speechbrain's spkrec-xvect-voxceleb extracts from a
49
+ reference audio file - what SpeechT5 conditions its voice on.
50
+
51
+ Args:
52
+ location: Path to a reference audio file (an 'asset:' reference has
53
+ already resolved to this by the time a task sees it).
54
+ device: Where to run the encoder.
55
+ dtype: The speech pipeline's own dtype - the embedding is concatenated
56
+ with the pipe's hidden states inside the decoder prenet, so it must
57
+ match the pipe's precision (fp16 on CUDA) or the linear layer there
58
+ refuses a Float x Half product.
59
+ """
60
+ from speechbrain.inference.speaker import EncoderClassifier
61
+
62
+ def load_encoder():
63
+ return EncoderClassifier.from_hparams(
64
+ source=_SPEAKER_ENCODER_MODEL,
65
+ run_opts={"device": str(device)},
66
+ )
67
+
68
+ encoder = cached_model(
69
+ ("speaker_encoder", _SPEAKER_ENCODER_MODEL, str(device)), load_encoder
70
+ )
71
+
72
+ waveform, sample_rate = load_audio(location)
73
+ waveform = resample_waveform(waveform, sample_rate, _SPEAKER_ENCODER_SAMPLE_RATE)
74
+ mono = waveform.mean(axis=0)
75
+
76
+ with torch.no_grad():
77
+ embedding = encoder.encode_batch(torch.as_tensor(mono).unsqueeze(0))
78
+ embedding = torch.nn.functional.normalize(embedding, dim=2)
79
+ # SpeechT5's generate() wants (batch, 512); the encoder's raw output is
80
+ # (1, 1, 512), so squeeze collapses it back to (512,) before restoring
81
+ # the batch dimension the model actually requires
82
+ return embedding.squeeze().unsqueeze(0).to(device=device, dtype=dtype)
83
+
84
+
85
+ def generate_speech(text=None, device="cpu", seed=None, **kwargs):
86
+ """Speak a line of text with a local text-to-speech model.
87
+
88
+ Args:
89
+ text: The line to speak. Mutually exclusive with messages - exactly
90
+ one of the two is required.
91
+ device: Target device ("cuda", "mps", "cpu").
92
+ seed: The workflow/step-resolved seed, when one was set. A pipeline
93
+ step gets an explicit torch.Generator; transformers' generate()
94
+ takes no such argument, so reproducing Bark (and any other
95
+ sampling model here) means seeding the global RNG immediately
96
+ before the call - unseeded when None, matching every other
97
+ unseeded task (#261).
98
+ **kwargs:
99
+ messages: Chat-templated input for a model such as VibeVoice that
100
+ takes a conversation rather than a bare string - a list of
101
+ {"role": ..., "content": ...} dicts. Passed straight through
102
+ as the pipeline's text_inputs, which applies the model's own
103
+ chat template; a model with no chat template configured (Bark
104
+ and friends) raises when handed this instead of text.
105
+ model_name: HuggingFace model ID. Defaults to suno/bark-small.
106
+ voice_preset: The speaker to use, for a model that has presets -
107
+ "v2/en_speaker_6" and friends for Bark. This is a preprocessing
108
+ argument: it selects the speaker before generation rather than
109
+ parameterizing it, which is why it is named here rather than left
110
+ to forward_params, where it would be silently dropped.
111
+ speaker_embedding: Path to a reference audio file (typically an
112
+ 'asset:' reference) whose voice a SpeechT5 model should speak
113
+ in. Reduced to an x-vector with speechbrain's
114
+ spkrec-xvect-voxceleb and injected into forward_params as
115
+ 'speaker_embeddings' - SpeechT5 is the only pipeline here that
116
+ conditions on one, and it is not optional for SpeechT5: the
117
+ model refuses to generate without a speaker embedding, so this
118
+ is required whenever model_name is a SpeechT5 checkpoint. A
119
+ VITS model's speaker instead takes a plain 'speaker_id' int,
120
+ which already reaches the model unchanged through
121
+ forward_params and needs no argument of its own.
122
+ forward_params: Passed to the model's forward/generate call.
123
+ generate_kwargs: Ad-hoc generation settings for a generative model -
124
+ temperature, do_sample and so on.
125
+
126
+ Returns:
127
+ An AudioTrack holding the waveform, shaped (channels, samples), and the
128
+ sample rate the model generated it at.
129
+
130
+ Raises:
131
+ ValueError: If the model reports no sample rate for what it generated,
132
+ or if text/messages are both given or neither is.
133
+ """
134
+ messages = kwargs.get("messages", None)
135
+ if (text is None) == (messages is None):
136
+ raise ValueError(
137
+ "generate_speech needs exactly one of 'text' or 'messages' - "
138
+ "a plain line to speak, or chat-templated input for a model "
139
+ "such as VibeVoice"
140
+ )
141
+ if messages is not None:
142
+ valid = isinstance(messages, list) and len(messages) > 0
143
+ if valid:
144
+ for message in messages:
145
+ if (
146
+ not isinstance(message, dict)
147
+ or not isinstance(message.get("role"), str)
148
+ or not isinstance(message.get("content"), str)
149
+ ):
150
+ valid = False
151
+ break
152
+ if not valid:
153
+ raise ValueError(
154
+ "generate_speech's 'messages' needs a non-empty list of "
155
+ "{'role': ..., 'content': ...} dicts, both strings - a bare "
156
+ "string or a wrong-keyed dict is not chat-templated input"
157
+ )
158
+ text_inputs = messages if messages is not None else text
159
+
160
+ model_name = kwargs.get("model_name", _DEFAULT_MODEL)
161
+ voice_preset = kwargs.get("voice_preset", None)
162
+ speaker_embedding = kwargs.get("speaker_embedding", None)
163
+ dtype = preferred_task_dtype(device)
164
+
165
+ def load_pipe():
166
+ logger.info(f"Generating speech with {model_name} on {device}")
167
+ placement = hf_pipeline_placement(device)
168
+ return hf_pipeline(
169
+ "text-to-speech",
170
+ model=model_name,
171
+ torch_dtype=dtype,
172
+ **placement,
173
+ )
174
+
175
+ pipe = cached_model(
176
+ ("speech_generation", model_name, str(device), str(dtype)),
177
+ load_pipe,
178
+ )
179
+
180
+ if voice_preset and getattr(pipe, "processor", None) is None:
181
+ # A single-voice model has no processor to hand the preset to;
182
+ # transformers logs the kwarg as unrecognised and speaks anyway
183
+ raise ValueError(
184
+ f"{model_name} takes no 'voice_preset' - it has one voice. Drop the "
185
+ "preset, or use a model with speaker presets such as suno/bark-small"
186
+ )
187
+
188
+ forward_params = kwargs.get("forward_params") or {}
189
+ if speaker_embedding:
190
+ model_type = getattr(getattr(pipe.model, "config", None), "model_type", None)
191
+ if model_type != "speecht5":
192
+ # Only SpeechT5 conditions on an x-vector; a model that does not
193
+ # would drop 'speaker_embeddings' as an unrecognised forward kwarg
194
+ # and generate in its own voice, same failure mode as voice_preset
195
+ raise ValueError(
196
+ f"{model_name} takes no 'speaker_embedding' - only a SpeechT5 "
197
+ "model conditions on an x-vector. Drop it, or use a SpeechT5 "
198
+ "model such as microsoft/speecht5_tts"
199
+ )
200
+ forward_params = {
201
+ **forward_params,
202
+ "speaker_embeddings": _speaker_embedding_tensor(
203
+ speaker_embedding, device, dtype
204
+ ),
205
+ }
206
+
207
+ if messages is not None:
208
+ logger.info(f"Speaking {len(messages)} chat-templated message(s)")
209
+ else:
210
+ logger.info(f"Speaking: {text[:100]}{'...' if len(text) > 100 else ''}")
211
+ if seed is not None:
212
+ torch.manual_seed(seed)
213
+ output = pipe(
214
+ text_inputs,
215
+ preprocess_params={"voice_preset": voice_preset} if voice_preset else {},
216
+ forward_params=forward_params,
217
+ generate_kwargs=kwargs.get("generate_kwargs") or {},
218
+ )
219
+
220
+ sample_rate = output.get("sampling_rate")
221
+ if sample_rate is None:
222
+ # Saving at the 44100 default instead would be quietly wrong rather than
223
+ # loud, and speech at the wrong rate is wrong in pitch as well as length
224
+ raise ValueError(
225
+ f"{model_name} reported no sample rate for the speech it generated"
226
+ )
227
+
228
+ return AudioTrack(as_channels_samples(output["audio"]), int(sample_rate))
dw/tasks/stabilize.py ADDED
@@ -0,0 +1,129 @@
1
+ """Hold a generated clip's framing still.
2
+
3
+ A video model drifts. Ask it for a locked camera and the framing still wanders
4
+ over a few seconds - a slow slide of the whole picture that nothing in the
5
+ prompt asked for. It goes unnoticed inside a single shot, and becomes obvious
6
+ the moment two shots are cut together and the subject jumps back to where it
7
+ started.
8
+
9
+ The wander is a global translation, so phase correlation measures it: the
10
+ cross-power spectrum of consecutive frames peaks at the shift between them.
11
+ Accumulating those shifts gives the clip's drift, and shifting every frame back
12
+ by it puts the framing where it began.
13
+ """
14
+
15
+ import logging
16
+
17
+ import numpy as np
18
+ from PIL import Image
19
+
20
+ from ..result import AudioVideo
21
+ from ..shots import carried_shots
22
+ from .video_utils import frames_as_pil_list, load_audio_video
23
+
24
+ logger = logging.getLogger("dw")
25
+
26
+
27
+ def _pair_shift(previous, following, window):
28
+ """The translation between two grayscale frames, in whole pixels."""
29
+ height, width = previous.shape
30
+ a = np.fft.rfft2(previous * window)
31
+ b = np.fft.rfft2(following * window)
32
+ cross = a * np.conj(b)
33
+ cross /= np.abs(cross) + 1e-8
34
+ correlation = np.fft.irfft2(cross, s=previous.shape)
35
+
36
+ # The peak of this cross-power spectrum sits at the negative of the
37
+ # displacement, so it is negated here and every caller reads (dx, dy) as
38
+ # "how far the picture moved between these two frames".
39
+ peak = np.unravel_index(np.argmax(correlation), correlation.shape)
40
+ dy = peak[0] - height if peak[0] > height // 2 else peak[0]
41
+ dx = peak[1] - width if peak[1] > width // 2 else peak[1]
42
+ return -dx, -dy
43
+
44
+
45
+ def _moving_average(values, window):
46
+ """Trajectory smoothed over `window` frames, with the ends held."""
47
+ pad = window // 2
48
+ padded = np.pad(values, ((pad, pad), (0, 0)), mode="edge")
49
+ kernel = np.ones(window, dtype=np.float32) / window
50
+ return np.stack(
51
+ [
52
+ np.convolve(padded[:, axis], kernel, mode="valid")[: len(values)]
53
+ for axis in (0, 1)
54
+ ],
55
+ axis=1,
56
+ )
57
+
58
+
59
+ def stabilize_video(clip, smooth=0):
60
+ """Task command: remove a generated clip's accumulated framing drift.
61
+
62
+ Args:
63
+ clip: The video - a frame list, a frame array or tensor, an
64
+ AudioVideo, whose soundtrack is carried through untouched, or the
65
+ path or URL of a video file, which is read with its audio, so a
66
+ shot an earlier run wrote can be steadied without regenerating it.
67
+ The argument is deliberately not called "video": the engine loads
68
+ an argument by that name itself, as bare frames, which would strip
69
+ the soundtrack off before this ever saw it
70
+ smooth: 0 locks the framing to the first frame, which is what a shot
71
+ generated from a pinned keyframe wants. A window in frames instead
72
+ removes only the wander faster than that window, so a slow
73
+ deliberate move survives and the drift around it does not
74
+ Returns:
75
+ The stabilized clip, cropped to the region every frame covers and
76
+ resized back to its original size - an AudioVideo when one came in
77
+ """
78
+ if isinstance(clip, str):
79
+ clip = load_audio_video(clip)
80
+
81
+ frames = frames_as_pil_list(clip)
82
+ if len(frames) < 2:
83
+ return clip
84
+
85
+ width, height = frames[0].size
86
+ gray = [np.asarray(f.convert("L"), dtype=np.float32) for f in frames]
87
+ window = np.outer(np.hanning(height), np.hanning(width))
88
+
89
+ trajectory = np.zeros((len(frames), 2), dtype=np.float32)
90
+ for index in range(1, len(frames)):
91
+ dx, dy = _pair_shift(gray[index - 1], gray[index], window)
92
+ trajectory[index] = trajectory[index - 1] + (dx, dy)
93
+
94
+ target = _moving_average(trajectory, smooth) if smooth > 1 else 0.0
95
+ correction = np.rint(trajectory - target).astype(np.int32)
96
+ logger.debug(
97
+ f"Stabilizing {len(frames)} frames, peak drift "
98
+ f"{np.abs(correction).max()}px of {max(width, height)}"
99
+ )
100
+ if not correction.any():
101
+ return clip
102
+
103
+ # Shifting frames back uncovers their edges, so keep only the rectangle
104
+ # every frame still covers, then put that back at the original size.
105
+ shifts = -correction
106
+ left = int(max(0, shifts[:, 0].max()))
107
+ right = int(width + min(0, shifts[:, 0].min()))
108
+ top = int(max(0, shifts[:, 1].max()))
109
+ bottom = int(height + min(0, shifts[:, 1].min()))
110
+
111
+ held = []
112
+ for frame, (sx, sy) in zip(frames, shifts):
113
+ moved = np.roll(np.asarray(frame), (int(sy), int(sx)), axis=(0, 1))
114
+ held.append(
115
+ Image.fromarray(moved[top:bottom, left:right]).resize(
116
+ (width, height), Image.LANCZOS
117
+ )
118
+ )
119
+
120
+ if isinstance(clip, AudioVideo):
121
+ # Same frames, one for one, so every shot boundary still holds
122
+ return AudioVideo(
123
+ held,
124
+ clip.audio,
125
+ clip.sample_rate,
126
+ fps=clip.fps,
127
+ shots=carried_shots(clip),
128
+ )
129
+ return held