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,266 @@
1
+ """Waveform utilities for audio tasks and segment-chained video generation.
2
+
3
+ Waveforms are handled as (channels, samples) float32 numpy arrays throughout -
4
+ as_channels_samples normalizes the shapes pipelines and files actually produce
5
+ into that layout.
6
+ """
7
+
8
+ import io
9
+ import logging
10
+
11
+ import numpy
12
+ import soundfile
13
+ import torch
14
+
15
+ from ..security import (
16
+ validate_path,
17
+ validate_url,
18
+ validate_file_extension,
19
+ ALLOWED_AUDIO_EXTENSIONS,
20
+ )
21
+
22
+ logger = logging.getLogger("dw")
23
+
24
+ # A few milliseconds of fade applied on each side of a butt-joined seam so the
25
+ # discontinuity does not click
26
+ DECLICK_MS = 3.0
27
+
28
+
29
+ def as_channels_samples(audio):
30
+ """Normalize a waveform to a (channels, samples) float32 numpy array.
31
+
32
+ Accepts torch tensors or numpy arrays shaped (samples,), (channels, samples),
33
+ (samples, channels), or a one-item batch (1, channels, samples). Channel
34
+ position is decided the way normalize_audio in result.py decides it: there
35
+ are always more samples than channels.
36
+ """
37
+ if torch.is_tensor(audio):
38
+ audio = audio.detach().cpu().float().numpy()
39
+ audio = numpy.asarray(audio, dtype=numpy.float32)
40
+
41
+ if audio.ndim == 1:
42
+ return audio[numpy.newaxis, :]
43
+
44
+ if audio.ndim == 3:
45
+ if audio.shape[0] != 1:
46
+ raise ValueError(f"Cannot normalize a waveform batch of {audio.shape[0]}")
47
+ audio = audio[0]
48
+
49
+ if audio.ndim != 2:
50
+ raise ValueError(f"A waveform must have 1-3 dimensions, got {audio.ndim}")
51
+
52
+ if audio.shape[0] > audio.shape[1]: # (samples, channels) -> transpose
53
+ audio = audio.T
54
+
55
+ return numpy.ascontiguousarray(audio)
56
+
57
+
58
+ def frames_to_samples(frames, fps, sample_rate):
59
+ """The number of audio samples spanning a run of video frames."""
60
+ return int(round(frames / fps * sample_rate))
61
+
62
+
63
+ def slice_samples(waveform, start, length):
64
+ """Cut length samples out of a (channels, samples) waveform from start.
65
+
66
+ A slice reaching past the end of the waveform is zero-padded to the
67
+ requested length, so frame-aligned slicing near the end of a track always
68
+ yields full-size chunks.
69
+ """
70
+ channels, total = waveform.shape
71
+ piece = waveform[:, start : start + length]
72
+ if piece.shape[1] < length:
73
+ padding = numpy.zeros((channels, length - piece.shape[1]), dtype=waveform.dtype)
74
+ piece = numpy.concatenate([piece, padding], axis=1)
75
+ return piece
76
+
77
+
78
+ def equal_power_crossfade_join(previous, head, following, sample_rate, crossfade_ms):
79
+ """Join two segments' audio at a seam without changing the total duration.
80
+
81
+ previous ends at the seam. head is the audio trimmed off the next segment's
82
+ start - it covers the same stretch of time as the tail of previous, so the
83
+ two are blended with an equal-power crossfade over the last
84
+ min(crossfade_ms, len(head)) of that stretch. following is the next
85
+ segment's on-timeline audio and is appended unchanged.
86
+
87
+ With no head material (nothing was trimmed), the seam gets a short declick
88
+ ramp instead - a few milliseconds of fade-out and fade-in in place.
89
+ """
90
+ previous, head, following = _matched_channels(previous, head, following)
91
+
92
+ window = min(
93
+ int(crossfade_ms / 1000.0 * sample_rate),
94
+ head.shape[1],
95
+ previous.shape[1],
96
+ )
97
+
98
+ if window == 0:
99
+ return _declick_join(previous, following, sample_rate)
100
+
101
+ fade_out, fade_in = _equal_power_ramps(window)
102
+ blended = previous[:, -window:] * fade_out + head[:, -window:] * fade_in
103
+ return numpy.concatenate([previous[:, :-window], blended, following], axis=1)
104
+
105
+
106
+ def crossfade_concat(waveforms, sample_rate, crossfade_ms):
107
+ """Concatenate waveforms, overlapping each seam by an equal-power crossfade.
108
+
109
+ The classic crossfade: each seam overlaps the two waveforms by the fade
110
+ window, so the result is shorter than the plain sum by one window per seam.
111
+ """
112
+ waveforms = [as_channels_samples(waveform) for waveform in waveforms]
113
+ if not waveforms:
114
+ raise ValueError("No waveforms to concatenate")
115
+
116
+ result = waveforms[0]
117
+ for following in waveforms[1:]:
118
+ result, following = _matched_channels(result, following)
119
+ window = min(
120
+ int(crossfade_ms / 1000.0 * sample_rate),
121
+ result.shape[1],
122
+ following.shape[1],
123
+ )
124
+ if window == 0:
125
+ result = _declick_join(result, following, sample_rate)
126
+ continue
127
+
128
+ fade_out, fade_in = _equal_power_ramps(window)
129
+ blended = result[:, -window:] * fade_out + following[:, :window] * fade_in
130
+ result = numpy.concatenate(
131
+ [result[:, :-window], blended, following[:, window:]], axis=1
132
+ )
133
+
134
+ return result
135
+
136
+
137
+ def load_audio(location, base_dir=None):
138
+ """Load an audio file from a local path or http(s) URL.
139
+
140
+ Returns:
141
+ Tuple of a (channels, samples) float32 waveform and its sample rate
142
+ """
143
+ if location.startswith(("http://", "https://")):
144
+ import requests
145
+
146
+ validated_url = validate_url(location)
147
+ logger.debug(f"Downloading audio from {validated_url}")
148
+ response = requests.get(validated_url, timeout=60)
149
+ response.raise_for_status()
150
+ data, sample_rate = soundfile.read(
151
+ io.BytesIO(response.content), dtype="float32"
152
+ )
153
+ else:
154
+ validated_path = validate_path(location, base_dir=base_dir, allow_create=False)
155
+ validate_file_extension(validated_path, ALLOWED_AUDIO_EXTENSIONS)
156
+ logger.debug(f"Reading audio from {validated_path}")
157
+ data, sample_rate = soundfile.read(validated_path, dtype="float32")
158
+
159
+ # soundfile returns (samples,) or (samples, channels)
160
+ return as_channels_samples(data), sample_rate
161
+
162
+
163
+ def slice_audio(
164
+ audio,
165
+ start_seconds=None,
166
+ duration_seconds=None,
167
+ start_frame=None,
168
+ num_frames=None,
169
+ fps=None,
170
+ sample_rate=None,
171
+ ):
172
+ """Task command: cut a slice out of an audio track.
173
+
174
+ The slice is addressed either in seconds (start_seconds + duration_seconds)
175
+ or in video frames (start_frame + num_frames + fps). Slices reaching past
176
+ the end of the track are zero-padded.
177
+
178
+ Args:
179
+ audio: Path or URL of an audio file, or a waveform (which needs
180
+ sample_rate alongside it)
181
+ sample_rate: Sample rate of a waveform passed directly; ignored for
182
+ files, which carry their own
183
+
184
+ Returns:
185
+ The slice as a (samples, channels) float32 array - the layout audio
186
+ results are saved in
187
+ """
188
+ if isinstance(audio, str):
189
+ waveform, sample_rate = load_audio(audio)
190
+ else:
191
+ if sample_rate is None:
192
+ raise ValueError("slice_audio needs 'sample_rate' with a raw waveform")
193
+ waveform = as_channels_samples(audio)
194
+
195
+ if start_seconds is not None or duration_seconds is not None:
196
+ if start_seconds is None or duration_seconds is None:
197
+ raise ValueError(
198
+ "slice_audio needs both 'start_seconds' and 'duration_seconds'"
199
+ )
200
+ start = int(round(start_seconds * sample_rate))
201
+ length = int(round(duration_seconds * sample_rate))
202
+ elif start_frame is not None or num_frames is not None:
203
+ if start_frame is None or num_frames is None or fps is None:
204
+ raise ValueError(
205
+ "slice_audio needs 'start_frame', 'num_frames' and 'fps' together"
206
+ )
207
+ start = frames_to_samples(start_frame, fps, sample_rate)
208
+ length = frames_to_samples(num_frames, fps, sample_rate)
209
+ else:
210
+ raise ValueError(
211
+ "slice_audio needs either 'start_seconds'/'duration_seconds' or "
212
+ "'start_frame'/'num_frames'/'fps'"
213
+ )
214
+
215
+ return slice_samples(waveform, start, length).T
216
+
217
+
218
+ def crossfade_audio(audios, crossfade_ms=75, sample_rate=None):
219
+ """Task command: join audio tracks with an equal-power crossfade.
220
+
221
+ Each seam overlaps the two tracks by the fade window, so the result is
222
+ shorter than the plain sum by one window per seam.
223
+
224
+ Args:
225
+ audios: The waveforms to join, in order
226
+ crossfade_ms: Length of each crossfade
227
+ sample_rate: Sample rate of the waveforms
228
+
229
+ Returns:
230
+ The joined track as a (samples, channels) float32 array
231
+ """
232
+ if sample_rate is None:
233
+ raise ValueError("crossfade_audio needs 'sample_rate'")
234
+ return crossfade_concat(audios, sample_rate, crossfade_ms).T
235
+
236
+
237
+ def _equal_power_ramps(window):
238
+ """Cosine/sine fade curves that sum to constant power across the window."""
239
+ theta = numpy.linspace(0.0, numpy.pi / 2.0, window, endpoint=False)
240
+ return numpy.cos(theta, dtype=numpy.float32), numpy.sin(theta, dtype=numpy.float32)
241
+
242
+
243
+ def _declick_join(previous, following, sample_rate):
244
+ """Butt-join two waveforms with a short fade on each side of the seam."""
245
+ ramp = int(DECLICK_MS / 1000.0 * sample_rate)
246
+ ramp = min(ramp, previous.shape[1], following.shape[1])
247
+ if ramp > 0:
248
+ fade_out, fade_in = _equal_power_ramps(ramp)
249
+ previous = previous.copy()
250
+ following = following.copy()
251
+ previous[:, -ramp:] *= fade_out # cos: 1 down to ~0
252
+ following[:, :ramp] *= fade_in # sin: ~0 up to 1
253
+ return numpy.concatenate([previous, following], axis=1)
254
+
255
+
256
+ def _matched_channels(*waveforms):
257
+ """Tile mono up so every waveform has the same channel count."""
258
+ channels = max(waveform.shape[0] for waveform in waveforms)
259
+ return tuple(
260
+ (
261
+ numpy.tile(waveform, (channels, 1))
262
+ if waveform.shape[0] == 1 and channels > 1
263
+ else waveform
264
+ )
265
+ for waveform in waveforms
266
+ )
@@ -0,0 +1,43 @@
1
+ from PIL import Image
2
+ import torch
3
+ from torchvision import transforms
4
+ from transformers import AutoModelForImageSegmentation
5
+
6
+ from .model_cache import cached_model
7
+
8
+ _MODEL_NAME = "briaai/RMBG-2.0"
9
+
10
+
11
+ def remove_background(image: Image, device) -> Image:
12
+ # Model settings
13
+ def load_model():
14
+ model = AutoModelForImageSegmentation.from_pretrained(
15
+ _MODEL_NAME, trust_remote_code=True
16
+ )
17
+ model.to(device)
18
+ model.eval()
19
+ return model
20
+
21
+ model = cached_model(("background_remover", _MODEL_NAME, str(device)), load_model)
22
+
23
+ # Data settings
24
+ transform_image = transforms.Compose(
25
+ [
26
+ transforms.Resize((1024, 1024)),
27
+ transforms.ToTensor(),
28
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
29
+ ]
30
+ )
31
+
32
+ working_copy = image.copy()
33
+ input_images = transform_image(working_copy).unsqueeze(0).to(device)
34
+
35
+ # Prediction
36
+ with torch.no_grad():
37
+ preds = model(input_images)[-1].sigmoid().cpu()
38
+ pred = preds[0].squeeze()
39
+ pred_pil = transforms.ToPILImage()(pred)
40
+ mask = pred_pil.resize(working_copy.size)
41
+ working_copy.putalpha(mask)
42
+
43
+ return working_copy
dw/tasks/borders.py ADDED
@@ -0,0 +1,113 @@
1
+ from PIL import Image
2
+
3
+
4
+ def add_border_and_mask(
5
+ image, zoom_all=1.0, zoom_left=0, zoom_right=0, zoom_up=0, zoom_down=0, overlap=0
6
+ ):
7
+ """Adds a black border around the image with individual side control and mask overlap"""
8
+ orig_width, orig_height = image.size
9
+
10
+ # Calculate padding for each side (in pixels)
11
+ left_pad = int(orig_width * zoom_left)
12
+ right_pad = int(orig_width * zoom_right)
13
+ top_pad = int(orig_height * zoom_up)
14
+ bottom_pad = int(orig_height * zoom_down)
15
+
16
+ # Calculate overlap in pixels
17
+ overlap_left = int(orig_width * overlap)
18
+ overlap_right = int(orig_width * overlap)
19
+ overlap_top = int(orig_height * overlap)
20
+ overlap_bottom = int(orig_height * overlap)
21
+
22
+ # If using the all-sides zoom, add it to each side
23
+ if zoom_all > 1.0:
24
+ extra_each_side = (zoom_all - 1.0) / 2
25
+ left_pad += int(orig_width * extra_each_side)
26
+ right_pad += int(orig_width * extra_each_side)
27
+ top_pad += int(orig_height * extra_each_side)
28
+ bottom_pad += int(orig_height * extra_each_side)
29
+
30
+ # Calculate new dimensions (ensure they're multiples of 32)
31
+ new_width = 32 * round((orig_width + left_pad + right_pad) / 32)
32
+ new_height = 32 * round((orig_height + top_pad + bottom_pad) / 32)
33
+
34
+ # Create new image with black border
35
+ bordered_image = Image.new("RGB", (new_width, new_height), (0, 0, 0))
36
+ # Paste original image in position
37
+ paste_x = left_pad
38
+ paste_y = top_pad
39
+ bordered_image.paste(image, (paste_x, paste_y))
40
+
41
+ # Create mask (white where the border is, black where the original image was)
42
+ mask = Image.new("L", (new_width, new_height), 255) # White background
43
+ # Paste black rectangle with overlap adjustment
44
+ mask.paste(
45
+ 0,
46
+ (
47
+ paste_x + overlap_left, # Left edge moves right
48
+ paste_y + overlap_top, # Top edge moves down
49
+ paste_x + orig_width - overlap_right, # Right edge moves left
50
+ paste_y + orig_height - overlap_bottom, # Bottom edge moves up
51
+ ),
52
+ )
53
+
54
+ return {"bordered_image": bordered_image, "mask": mask}
55
+
56
+
57
+ def add_border_and_mask_with_size(image, width, height, overlap=0):
58
+ """
59
+ Resizes the original image to fit within the target dimensions while maintaining
60
+ its aspect ratio, then adds borders as needed to reach the exact target size.
61
+
62
+ Args:
63
+ image: PIL Image object
64
+ width: Target width in pixels
65
+ height: Target height in pixels
66
+ overlap: Mask overlap parameter (0-1 range)
67
+
68
+ Returns:
69
+ Dictionary with 'bordered_image' and 'mask'
70
+ """
71
+ # Ensure width and height are multiples of 32
72
+ width = 32 * round(width / 32)
73
+ height = 32 * round(height / 32)
74
+
75
+ # Get original dimensions
76
+ orig_width, orig_height = image.size
77
+ orig_aspect = orig_width / orig_height
78
+ target_aspect = width / height
79
+
80
+ # Resize image to fit within target dimensions while maintaining aspect ratio
81
+ if orig_aspect > target_aspect:
82
+ # Original is wider than target - fit width
83
+ new_width = width
84
+ new_height = int(width / orig_aspect)
85
+ resized_image = image.resize((new_width, new_height), Image.LANCZOS)
86
+ else:
87
+ # Original is taller than target - fit height
88
+ new_height = height
89
+ new_width = int(height * orig_aspect)
90
+ resized_image = image.resize((new_width, new_height), Image.LANCZOS)
91
+
92
+ # Now calculate padding to reach target dimensions
93
+ left_pad = (width - new_width) // 2
94
+ right_pad = width - new_width - left_pad
95
+ top_pad = (height - new_height) // 2
96
+ bottom_pad = height - new_height - top_pad
97
+
98
+ # Convert padding to zoom factors (relative to resized dimensions)
99
+ zoom_left = left_pad / new_width if new_width > 0 else 0
100
+ zoom_right = right_pad / new_width if new_width > 0 else 0
101
+ zoom_up = top_pad / new_height if new_height > 0 else 0
102
+ zoom_down = bottom_pad / new_height if new_height > 0 else 0
103
+
104
+ # Call the original function with calculated zoom parameters
105
+ return add_border_and_mask(
106
+ resized_image,
107
+ zoom_all=1.0,
108
+ zoom_left=zoom_left,
109
+ zoom_right=zoom_right,
110
+ zoom_up=zoom_up,
111
+ zoom_down=zoom_down,
112
+ overlap=overlap,
113
+ )
@@ -0,0 +1,80 @@
1
+ """Concatenate videos - and the audio generated with them - into one video.
2
+
3
+ The standalone counterpart of what a chained pipeline step does internally:
4
+ frames are joined end to end with an optional head trim on every video after
5
+ the first, and audio tracks are joined at each seam with an equal-power
6
+ crossfade drawn from the trimmed-off material, so video and audio stay in
7
+ sync.
8
+ """
9
+
10
+ import logging
11
+
12
+ from ..result import AudioVideo
13
+ from .audio_utils import (
14
+ as_channels_samples,
15
+ equal_power_crossfade_join,
16
+ frames_to_samples,
17
+ )
18
+ from .video_utils import frames_as_pil_list
19
+
20
+ logger = logging.getLogger("dw")
21
+
22
+
23
+ def concat_videos(videos, trim_frames=0, crossfade_ms=75, fps=None):
24
+ """Concatenate a list of videos into a single AudioVideo.
25
+
26
+ Args:
27
+ videos: The videos to join, in order - frame lists, frame arrays, or
28
+ AudioVideos (from gather_videos or previous_result references)
29
+ trim_frames: Frames dropped from the head of every video after the
30
+ first - the trim used when each video was generated from the
31
+ previous one's last frame
32
+ crossfade_ms: Equal-power crossfade at each audio seam, clamped to
33
+ the trimmed material
34
+ fps: Frame rate of the videos - required to join audio when trimming
35
+
36
+ Returns:
37
+ One AudioVideo; its audio is None when no input video carries any
38
+ """
39
+ if not isinstance(videos, list) or not videos:
40
+ raise ValueError("concat_videos needs a non-empty list of videos")
41
+
42
+ frames = []
43
+ audio = None
44
+ sample_rate = None
45
+
46
+ for index, video in enumerate(videos):
47
+ head_trim = trim_frames if index > 0 else 0
48
+ frames.extend(frames_as_pil_list(video)[head_trim:])
49
+
50
+ if not isinstance(video, AudioVideo) or video.audio is None:
51
+ continue
52
+
53
+ waveform = as_channels_samples(video.audio)
54
+ if audio is None:
55
+ audio, sample_rate = waveform, video.sample_rate
56
+ continue
57
+
58
+ if video.sample_rate != sample_rate:
59
+ raise ValueError(
60
+ f"Videos carry audio at different sample rates: "
61
+ f"{sample_rate} then {video.sample_rate}"
62
+ )
63
+ if head_trim > 0 and fps is None:
64
+ raise ValueError(
65
+ "concat_videos needs 'fps' to trim audio in step with the frames"
66
+ )
67
+
68
+ trim_samples = (
69
+ frames_to_samples(head_trim, fps, sample_rate) if head_trim else 0
70
+ )
71
+ audio = equal_power_crossfade_join(
72
+ audio,
73
+ waveform[:, :trim_samples],
74
+ waveform[:, trim_samples:],
75
+ sample_rate,
76
+ crossfade_ms,
77
+ )
78
+
79
+ logger.debug(f"Concatenated {len(videos)} videos into {len(frames)} frames")
80
+ return AudioVideo(frames, audio, sample_rate)
@@ -0,0 +1,54 @@
1
+ import torch
2
+ import numpy as np
3
+ from transformers import pipeline
4
+ from torchvision import transforms
5
+
6
+ from .. import preferred_task_dtype
7
+ from .model_cache import cached_model
8
+
9
+
10
+ def make_hint_tensor(image, device, dtype=None):
11
+ """Estimate depth and return it as a hint tensor for a controlnet pipeline.
12
+
13
+ Args:
14
+ image: Image to estimate depth from
15
+ device: Device to run the estimator on and place the hint on
16
+ dtype: Dtype of the hint, defaulting to the one the device works best in.
17
+ The hint has to match the dtype of the pipeline that consumes it.
18
+
19
+ Returns:
20
+ Depth hint as a tensor of shape (1, 3, height, width)
21
+ """
22
+ depth_estimator = cached_model(
23
+ ("depth_estimator", str(device)),
24
+ lambda: pipeline("depth-estimation", device=device),
25
+ )
26
+
27
+ image = depth_estimator(image)["depth"]
28
+ image = np.array(image)
29
+ image = image[:, :, None]
30
+ image = np.concatenate([image, image, image], axis=2)
31
+ detected_map = torch.from_numpy(image).float() / 255.0
32
+ hint = detected_map.permute(2, 0, 1)
33
+
34
+ if dtype is None:
35
+ dtype = preferred_task_dtype(device)
36
+
37
+ return hint.unsqueeze(0).to(device=device, dtype=dtype)
38
+
39
+
40
+ def make_hint_image(image, device, dtype=None):
41
+ """Estimate depth and return it as an image.
42
+
43
+ Args:
44
+ image: Image to estimate depth from
45
+ device: Device to run the estimator on
46
+ dtype: Dtype to compute the hint in - see make_hint_tensor
47
+
48
+ Returns:
49
+ Depth map as a PIL image
50
+ """
51
+ hint = make_hint_tensor(image, device, dtype)
52
+ # Convert the tensor to a Pillow image
53
+ to_pil = transforms.ToPILImage()
54
+ return to_pil(hint[0].float().cpu())
@@ -0,0 +1,109 @@
1
+ """
2
+ Diffusion-based image upscaling via Stable Diffusion upscale pipelines.
3
+
4
+ Provides text-guided upscaling with better detail recovery than
5
+ traditional super-resolution models, especially for faces and textures.
6
+
7
+ Supports two modes:
8
+ - "x4" (default): StableDiffusionUpscalePipeline (4x, stabilityai/stable-diffusion-x4-upscaler)
9
+ - "x2": StableDiffusionLatentUpscalePipeline (2x, stabilityai/sd-x2-latent-upscaler)
10
+ """
11
+
12
+ import logging
13
+ import torch
14
+ import diffusers
15
+ from .. import preferred_task_dtype
16
+ from .model_cache import cached_model
17
+
18
+ logger = logging.getLogger("dw")
19
+
20
+ _MODELS = {
21
+ "x4": {
22
+ "pipeline_class": "StableDiffusionUpscalePipeline",
23
+ "model_name": "stabilityai/stable-diffusion-x4-upscaler",
24
+ },
25
+ "x2": {
26
+ "pipeline_class": "StableDiffusionLatentUpscalePipeline",
27
+ "model_name": "stabilityai/sd-x2-latent-upscaler",
28
+ },
29
+ }
30
+
31
+
32
+ def diffusion_upscale(image, device="cpu", **kwargs):
33
+ """Upscale an image using a Stable Diffusion upscale pipeline.
34
+
35
+ Args:
36
+ image: PIL Image to upscale.
37
+ device: Target device ("cuda", "mps", "cpu").
38
+ **kwargs:
39
+ prompt: Text guidance for upscaling (default: "").
40
+ negative_prompt: Negative text guidance (default: None).
41
+ mode: "x4" or "x2" (default: "x4").
42
+ model_name: Override the default model for the selected mode.
43
+ num_inference_steps: Denoising steps (default: 25).
44
+ guidance_scale: Classifier-free guidance scale (default: 9.0).
45
+ noise_level: Noise level for x4 mode (default: 20, ignored for x2).
46
+
47
+ Returns:
48
+ PIL Image (upscaled).
49
+ """
50
+ mode = kwargs.get("mode", "x4")
51
+ if mode not in _MODELS:
52
+ raise ValueError(f"mode must be one of {sorted(_MODELS.keys())}, got '{mode}'")
53
+
54
+ config = _MODELS[mode]
55
+ model_name = kwargs.get("model_name", config["model_name"])
56
+ prompt = kwargs.get("prompt", "")
57
+ negative_prompt = kwargs.get("negative_prompt", None)
58
+ num_inference_steps = int(kwargs.get("num_inference_steps", 25))
59
+ guidance_scale = float(kwargs.get("guidance_scale", 9.0))
60
+ noise_level = int(kwargs.get("noise_level", 20))
61
+
62
+ pipeline_class = getattr(diffusers, config["pipeline_class"])
63
+
64
+ dtype = preferred_task_dtype(device)
65
+
66
+ def load_pipe():
67
+ logger.info(f"Loading {config['pipeline_class']} from {model_name} to {device}")
68
+ pipe = pipeline_class.from_pretrained(
69
+ model_name,
70
+ torch_dtype=dtype,
71
+ )
72
+ pipe.to(device)
73
+ return pipe
74
+
75
+ pipe = cached_model(
76
+ (
77
+ "diffusion_upscale",
78
+ config["pipeline_class"],
79
+ model_name,
80
+ str(device),
81
+ str(dtype),
82
+ ),
83
+ load_pipe,
84
+ )
85
+
86
+ call_kwargs = {
87
+ "prompt": prompt,
88
+ "image": image,
89
+ "num_inference_steps": num_inference_steps,
90
+ "guidance_scale": guidance_scale,
91
+ }
92
+
93
+ if negative_prompt is not None:
94
+ call_kwargs["negative_prompt"] = negative_prompt
95
+
96
+ if mode == "x4":
97
+ call_kwargs["noise_level"] = noise_level
98
+
99
+ logger.info(
100
+ f"Upscaling {image.width}x{image.height} with {mode} mode, "
101
+ f"{num_inference_steps} steps"
102
+ )
103
+
104
+ with torch.inference_mode():
105
+ result = pipe(**call_kwargs)
106
+
107
+ output = result.images[0]
108
+ logger.info(f"Upscaled to {output.width}x{output.height}")
109
+ return output
@@ -0,0 +1,24 @@
1
+ def format_chat_message(system_prompt, user_message):
2
+ return {
3
+ "text_inputs": [
4
+ {"role": "system", "content": system_prompt},
5
+ {
6
+ "role": "user",
7
+ "content": user_message,
8
+ },
9
+ ]
10
+ }
11
+
12
+
13
+ def batch_decode_post_process(processor, task, generated_ids):
14
+ generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
15
+
16
+ parsed_answer = processor.post_process_generation(generated_text, task=task)
17
+
18
+ return parsed_answer[task]
19
+
20
+
21
+ def get_dict_value(dict, key):
22
+ if key in dict:
23
+ return dict[key]
24
+ return None