diffusers-workflow 0.4.0a3__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
- diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
- diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
- diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
- dw/__init__.py +353 -0
- dw/arguments.py +906 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +145 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/events.py +78 -0
- dw/hub_cache.py +289 -0
- dw/introspection.py +458 -0
- dw/log_setup.py +45 -0
- dw/pipeline_processors/chain.py +750 -0
- dw/pipeline_processors/config_objects.py +235 -0
- dw/pipeline_processors/pipeline.py +1687 -0
- dw/pipeline_processors/remote.py +18 -0
- dw/previous_results.py +259 -0
- dw/prompt_weighting.py +378 -0
- dw/repl.py +298 -0
- dw/repl_commands.py +808 -0
- dw/repl_worker.py +129 -0
- dw/result.py +850 -0
- dw/run.py +92 -0
- dw/schema.py +24 -0
- dw/security.py +379 -0
- dw/serve.py +70 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +588 -0
- dw/server/jobs.py +547 -0
- dw/server/ui/assets/abap-08VXUWAP.js +1 -0
- dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
- dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
- dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
- dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
- dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
- dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
- dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
- dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
- dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
- dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
- dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
- dw/server/ui/assets/css-DIMkf-bt.js +3 -0
- dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
- dw/server/ui/assets/cssMode-CEh6hWi2.js +1 -0
- dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
- dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
- dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
- dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
- dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
- dw/server/ui/assets/editor.api-CExg3_mM.js +847 -0
- dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
- dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
- dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
- dw/server/ui/assets/freemarker2-DH6orYh2.js +3 -0
- dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
- dw/server/ui/assets/go-C-y9NEjX.js +1 -0
- dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
- dw/server/ui/assets/handlebars-CbrMVW4Q.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-YDNPZw2M.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
- dw/server/ui/assets/index-B7-VcYS-.css +1 -0
- dw/server/ui/assets/index-D_EiPU3b.js +13 -0
- dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
- dw/server/ui/assets/java-BEtHBSE6.js +1 -0
- dw/server/ui/assets/javascript-dYuBvioq.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-CUqLM39V.js +7 -0
- dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
- dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
- dw/server/ui/assets/less-B9JPFI3C.js +2 -0
- dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
- dw/server/ui/assets/liquid-D6vxBzMv.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.js +4 -0
- dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
- dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
- dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
- dw/server/ui/assets/mdx-SHQb6vmD.js +1 -0
- dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
- dw/server/ui/assets/monaco--ixms01u.css +1 -0
- dw/server/ui/assets/monaco-CP-s5rcP.js +56 -0
- dw/server/ui/assets/msdax-DauUninz.js +1 -0
- dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
- dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
- dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
- dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
- dw/server/ui/assets/perl-oz_6vUea.js +1 -0
- dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
- dw/server/ui/assets/php-nr791fC2.js +1 -0
- dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
- dw/server/ui/assets/postiats-43DmfD33.js +1 -0
- dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
- dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
- dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
- dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
- dw/server/ui/assets/python-x0_EGHq9.js +1 -0
- dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
- dw/server/ui/assets/r-BwWrilGY.js +1 -0
- dw/server/ui/assets/razor-BZC4LQDP.js +1 -0
- dw/server/ui/assets/redis-ClamHrr6.js +1 -0
- dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
- dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
- dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
- dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
- dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
- dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
- dw/server/ui/assets/scheme-BeGwcela.js +1 -0
- dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
- dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
- dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
- dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
- dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
- dw/server/ui/assets/sql-NEE52Syq.js +1 -0
- dw/server/ui/assets/st-DbInun42.js +1 -0
- dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
- dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
- dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
- dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
- dw/server/ui/assets/tsMode-BTfA6SbD.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-CWA4MsNk.js +1 -0
- dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
- dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
- dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
- dw/server/ui/assets/workers-CWU0uvj5.js +1 -0
- dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
- dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/settings.py +77 -0
- dw/step.py +132 -0
- dw/tasks/audio_utils.py +266 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/concat_videos.py +80 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +139 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +661 -0
- dw/tasks/interpolate_frames.py +227 -0
- dw/tasks/model_cache.py +39 -0
- dw/tasks/pair_audio.py +58 -0
- dw/tasks/qr_code.py +19 -0
- dw/tasks/restore_faces.py +175 -0
- dw/tasks/rife_model.py +192 -0
- dw/tasks/segment.py +121 -0
- dw/tasks/task.py +474 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +168 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +154 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +376 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +68 -0
- dw/validate.py +43 -0
- dw/variables.py +153 -0
- dw/worker.py +517 -0
- dw/workflow.py +553 -0
- dw/workflow_schema.json +1157 -0
- dw/workflows/augment_prompt.json +65 -0
- dw/workflows/describe_image.json +58 -0
- dw/workflows/h3_context_ir.json +57 -0
- dw/workflows/test.json +31 -0
dw/tasks/audio_utils.py
ADDED
|
@@ -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
|