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
|
@@ -0,0 +1,227 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Video frame interpolation via RIFE (Real-Time Intermediate Flow Estimation).
|
|
3
|
+
|
|
4
|
+
Takes a list of video frames and generates intermediate frames to increase
|
|
5
|
+
frame rate. Supports 2x, 4x, and 8x multipliers.
|
|
6
|
+
|
|
7
|
+
Model weights are downloaded from HuggingFace Hub on first use.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
import logging
|
|
11
|
+
import torch
|
|
12
|
+
|
|
13
|
+
from .tensor_image import pil_to_float_tensor as _pil_to_tensor, float_tensor_to_pil
|
|
14
|
+
|
|
15
|
+
logger = logging.getLogger("dw")
|
|
16
|
+
|
|
17
|
+
_VALID_MULTIPLIERS = {2, 4, 8}
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def interpolate_frames(video, device="cpu", **kwargs):
|
|
21
|
+
"""Interpolate between video frames using RIFE to increase frame rate.
|
|
22
|
+
|
|
23
|
+
Args:
|
|
24
|
+
video: List of PIL Images (video frames)
|
|
25
|
+
device: Target device ("cuda", "mps", "cpu")
|
|
26
|
+
**kwargs:
|
|
27
|
+
multiplier: Frame count multiplier — 2, 4, or 8 (default: 2)
|
|
28
|
+
model_name: HuggingFace repo with RIFE v4.13 weights (default: auto)
|
|
29
|
+
filename: Weights filename within the repo (default: auto)
|
|
30
|
+
|
|
31
|
+
Returns:
|
|
32
|
+
List of PIL Images with interpolated frames inserted.
|
|
33
|
+
"""
|
|
34
|
+
multiplier = int(kwargs.get("multiplier", 2))
|
|
35
|
+
model_name = kwargs.get("model_name", None)
|
|
36
|
+
filename = kwargs.get("filename", None)
|
|
37
|
+
|
|
38
|
+
if multiplier not in _VALID_MULTIPLIERS:
|
|
39
|
+
raise ValueError(
|
|
40
|
+
f"multiplier must be one of {sorted(_VALID_MULTIPLIERS)}, got {multiplier}"
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
if len(video) < 2:
|
|
44
|
+
raise ValueError(f"Need at least 2 frames to interpolate, got {len(video)}")
|
|
45
|
+
|
|
46
|
+
logger.info(
|
|
47
|
+
f"Interpolating {len(video)} frames with {multiplier}x multiplier on {device}"
|
|
48
|
+
)
|
|
49
|
+
|
|
50
|
+
model = _load_rife_model(device, model_name, filename)
|
|
51
|
+
|
|
52
|
+
passes = {2: 1, 4: 2, 8: 3}[multiplier]
|
|
53
|
+
frames = list(video)
|
|
54
|
+
|
|
55
|
+
for pass_num in range(passes):
|
|
56
|
+
logger.debug(
|
|
57
|
+
f"Interpolation pass {pass_num + 1}/{passes}: {len(frames)} frames"
|
|
58
|
+
)
|
|
59
|
+
frames = _interpolate_2x(frames, model)
|
|
60
|
+
|
|
61
|
+
logger.info(f"Interpolation complete: {len(video)} -> {len(frames)} frames")
|
|
62
|
+
return frames
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _interpolate_2x(frames, model):
|
|
66
|
+
"""Single pass of 2x interpolation — insert one frame between each pair."""
|
|
67
|
+
result = [frames[0]]
|
|
68
|
+
for i in range(len(frames) - 1):
|
|
69
|
+
mid_frame = model(frames[i], frames[i + 1])
|
|
70
|
+
result.append(mid_frame)
|
|
71
|
+
result.append(frames[i + 1])
|
|
72
|
+
return result
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
_DEFAULT_RIFE_REPO = "imaginairy/rife-interpolation"
|
|
76
|
+
_DEFAULT_RIFE_FILENAME = "rife-flownet-4.13.2.safetensors"
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def _pad_tensor(t, ph, pw):
|
|
80
|
+
"""Zero-pad a (1, 3, H, W) tensor on the bottom/right to (ph, pw)."""
|
|
81
|
+
h, w = t.shape[2], t.shape[3]
|
|
82
|
+
padding = (0, pw - w, 0, ph - h)
|
|
83
|
+
return torch.nn.functional.pad(t, padding)
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _padded_size(h, w, multiple=32):
|
|
87
|
+
"""Round (h, w) up to the next multiple (RIFE requires dims divisible by 32)."""
|
|
88
|
+
ph = ((h - 1) // multiple + 1) * multiple
|
|
89
|
+
pw = ((w - 1) // multiple + 1) * multiple
|
|
90
|
+
return ph, pw
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _make_flow_context(ph, pw, device):
|
|
94
|
+
"""Build the warp grid and flow divisors for one padded resolution."""
|
|
95
|
+
tenFlow_div = torch.tensor([(pw - 1.0) / 2.0, (ph - 1.0) / 2.0], device=device)
|
|
96
|
+
backwarp_tenGrid = torch.cat(
|
|
97
|
+
[
|
|
98
|
+
torch.linspace(-1.0, 1.0, pw, device=device)
|
|
99
|
+
.view(1, 1, 1, pw)
|
|
100
|
+
.expand(-1, -1, ph, -1),
|
|
101
|
+
torch.linspace(-1.0, 1.0, ph, device=device)
|
|
102
|
+
.view(1, 1, ph, 1)
|
|
103
|
+
.expand(-1, -1, -1, pw),
|
|
104
|
+
],
|
|
105
|
+
1,
|
|
106
|
+
)
|
|
107
|
+
timestep = torch.full((1, 1, ph, pw), 0.5, dtype=torch.float32, device=device)
|
|
108
|
+
return tenFlow_div, backwarp_tenGrid, timestep
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def _load_rife_model(device, model_name=None, filename=None):
|
|
112
|
+
"""Load RIFE model and return a callable that interpolates two frames.
|
|
113
|
+
|
|
114
|
+
Args:
|
|
115
|
+
device: Target device string ("cuda", "mps", "cpu").
|
|
116
|
+
model_name: Optional HuggingFace repo ID containing RIFE v4.13 weights.
|
|
117
|
+
Defaults to imaginairy/rife-interpolation.
|
|
118
|
+
filename: Optional weights filename within the repo. Defaults to
|
|
119
|
+
rife-flownet-4.13.2.safetensors. Both .safetensors and torch
|
|
120
|
+
checkpoint formats (.pkl/.pth) are supported.
|
|
121
|
+
|
|
122
|
+
Returns:
|
|
123
|
+
Callable that takes (frame1: PIL.Image, frame2: PIL.Image) -> PIL.Image
|
|
124
|
+
"""
|
|
125
|
+
from .rife_model import IFNet
|
|
126
|
+
from huggingface_hub import hf_hub_download
|
|
127
|
+
from .model_cache import cached_model
|
|
128
|
+
|
|
129
|
+
repo_id = model_name if model_name is not None else _DEFAULT_RIFE_REPO
|
|
130
|
+
weights_file = filename if filename is not None else _DEFAULT_RIFE_FILENAME
|
|
131
|
+
|
|
132
|
+
def load_net():
|
|
133
|
+
model_path = hf_hub_download(repo_id=repo_id, filename=weights_file)
|
|
134
|
+
|
|
135
|
+
logger.info(f"Loading RIFE IFNet v4.13 to {device}")
|
|
136
|
+
|
|
137
|
+
if weights_file.endswith(".safetensors"):
|
|
138
|
+
from safetensors.torch import load_file
|
|
139
|
+
|
|
140
|
+
state_dict = load_file(model_path)
|
|
141
|
+
else:
|
|
142
|
+
state_dict = torch.load(model_path, map_location="cpu", weights_only=True)
|
|
143
|
+
|
|
144
|
+
# Strip "module." prefix that comes from DataParallel-saved checkpoints
|
|
145
|
+
cleaned = {}
|
|
146
|
+
for k, v in state_dict.items():
|
|
147
|
+
cleaned[k.removeprefix("module.")] = v
|
|
148
|
+
|
|
149
|
+
net = IFNet()
|
|
150
|
+
net.load_state_dict(cleaned)
|
|
151
|
+
net.eval()
|
|
152
|
+
net.to(device)
|
|
153
|
+
return net
|
|
154
|
+
|
|
155
|
+
net = cached_model(
|
|
156
|
+
("interpolate_frames", repo_id, weights_file, str(device)), load_net
|
|
157
|
+
)
|
|
158
|
+
|
|
159
|
+
return _build_inference(net, device)
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
def _build_inference(net, device):
|
|
163
|
+
"""Build the (img1, img2) -> mid_frame callable for a loaded RIFE net.
|
|
164
|
+
|
|
165
|
+
All frames of a video share one padded resolution, so the flow-warp
|
|
166
|
+
context (tenFlow_div, backwarp grid, timestep) is computed once per
|
|
167
|
+
(padded_h, padded_w) and reused for every pair instead of being rebuilt
|
|
168
|
+
on every call.
|
|
169
|
+
|
|
170
|
+
`_interpolate_2x` walks frames pairwise: (f0, f1), (f1, f2), (f2, f3), ...
|
|
171
|
+
— the second frame of one pair is the same PIL object as the first frame
|
|
172
|
+
of the next. This closure carries the padded tensor it produced for a
|
|
173
|
+
pair's second frame forward, so that when that exact frame object shows
|
|
174
|
+
up again as a pair's first frame, its tensor is reused instead of being
|
|
175
|
+
re-derived from the PIL image (re-decoded, re-normalized, re-padded). A
|
|
176
|
+
cache miss (non-matching object, e.g. across separate interpolation
|
|
177
|
+
passes) simply falls back to a fresh conversion, so this is a pure
|
|
178
|
+
optimization with no behavioral effect — the reused tensor is bit-for-bit
|
|
179
|
+
what a fresh conversion of the same frame would produce.
|
|
180
|
+
"""
|
|
181
|
+
scale_list = [8, 4, 2, 1]
|
|
182
|
+
flow_context_cache = {}
|
|
183
|
+
carry = {"img": None}
|
|
184
|
+
|
|
185
|
+
def get_flow_context(ph, pw):
|
|
186
|
+
key = (ph, pw)
|
|
187
|
+
ctx = flow_context_cache.get(key)
|
|
188
|
+
if ctx is None:
|
|
189
|
+
ctx = _make_flow_context(ph, pw, device)
|
|
190
|
+
flow_context_cache[key] = ctx
|
|
191
|
+
return ctx
|
|
192
|
+
|
|
193
|
+
def inference(img1, img2):
|
|
194
|
+
"""Interpolate a single frame between two input frames."""
|
|
195
|
+
if img1 is carry["img"]:
|
|
196
|
+
t1_padded = carry["tensor"]
|
|
197
|
+
h, w, ph, pw = carry["h"], carry["w"], carry["ph"], carry["pw"]
|
|
198
|
+
else:
|
|
199
|
+
t1 = _pil_to_tensor(img1, device)
|
|
200
|
+
h, w = t1.shape[2], t1.shape[3]
|
|
201
|
+
ph, pw = _padded_size(h, w)
|
|
202
|
+
t1_padded = _pad_tensor(t1, ph, pw)
|
|
203
|
+
|
|
204
|
+
t2 = _pil_to_tensor(img2, device)
|
|
205
|
+
t2_padded = _pad_tensor(t2, ph, pw)
|
|
206
|
+
|
|
207
|
+
# Carry img2's padded tensor forward in case it's the next pair's img1.
|
|
208
|
+
carry["img"] = img2
|
|
209
|
+
carry["tensor"] = t2_padded
|
|
210
|
+
carry["h"], carry["w"], carry["ph"], carry["pw"] = h, w, ph, pw
|
|
211
|
+
|
|
212
|
+
tenFlow_div, backwarp_tenGrid, timestep = get_flow_context(ph, pw)
|
|
213
|
+
|
|
214
|
+
with torch.inference_mode():
|
|
215
|
+
_, _, merged = net(
|
|
216
|
+
t1_padded,
|
|
217
|
+
t2_padded,
|
|
218
|
+
timestep,
|
|
219
|
+
scale_list,
|
|
220
|
+
tenFlow_div,
|
|
221
|
+
backwarp_tenGrid,
|
|
222
|
+
)
|
|
223
|
+
|
|
224
|
+
mid = merged[3][:, :, :h, :w]
|
|
225
|
+
return float_tensor_to_pil(mid)
|
|
226
|
+
|
|
227
|
+
return inference
|
dw/tasks/model_cache.py
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
logger = logging.getLogger("dw")
|
|
4
|
+
|
|
5
|
+
# Loaded task models, keyed by whatever identifies a load - typically
|
|
6
|
+
# (task, model_name, device). step.py runs a task handler once per cartesian
|
|
7
|
+
# product iteration; without this, segmenting 20 images would load the same
|
|
8
|
+
# multi-gigabyte checkpoints 20 times over
|
|
9
|
+
_cache = {}
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def cached_model(key, factory):
|
|
13
|
+
"""Return the model for key, loading it with factory() on first use.
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
key: Hashable identity of the load - include the model name and the
|
|
17
|
+
device, plus anything else that changes what factory() builds
|
|
18
|
+
factory: Zero-argument callable performing the actual load
|
|
19
|
+
|
|
20
|
+
Returns:
|
|
21
|
+
The cached or freshly loaded model
|
|
22
|
+
"""
|
|
23
|
+
if key not in _cache:
|
|
24
|
+
logger.info(f"Loading task model: {key}")
|
|
25
|
+
_cache[key] = factory()
|
|
26
|
+
else:
|
|
27
|
+
logger.debug(f"Reusing cached task model: {key}")
|
|
28
|
+
return _cache[key]
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def clear_model_cache():
|
|
32
|
+
"""Release every cached task model.
|
|
33
|
+
|
|
34
|
+
Wired into the worker's memory cleanup - dropping the references here is
|
|
35
|
+
what lets gc and the allocator actually reclaim the weights.
|
|
36
|
+
"""
|
|
37
|
+
if _cache:
|
|
38
|
+
logger.info(f"Clearing {len(_cache)} cached task models")
|
|
39
|
+
_cache.clear()
|
dw/tasks/pair_audio.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
"""Pair a video with an audio track so the two are saved as one file.
|
|
2
|
+
|
|
3
|
+
A pipeline that generates its own soundtrack returns the pair together, and the
|
|
4
|
+
result muxes them into a single mp4. Anything that works on the frames alone -
|
|
5
|
+
a latent upsampler, an interpolator, an upscaler - returns frames without it, so
|
|
6
|
+
the soundtrack has to be carried across the step that dropped it. That is what
|
|
7
|
+
this does: it puts the two back together for the step that saves them.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
import logging
|
|
11
|
+
|
|
12
|
+
from ..result import AudioVideo
|
|
13
|
+
from .audio_utils import as_channels_samples
|
|
14
|
+
|
|
15
|
+
logger = logging.getLogger("dw")
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def pair_audio(video, audio, sample_rate=None):
|
|
19
|
+
"""Pair a video's frames with an audio track.
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
video: The frames - a frame list, a frame array or tensor, or an
|
|
23
|
+
AudioVideo whose own soundtrack is replaced by this one
|
|
24
|
+
audio: The soundtrack - a waveform, or an AudioVideo (or any object
|
|
25
|
+
carrying '.audio') to take it from, which brings its sample rate
|
|
26
|
+
along with it
|
|
27
|
+
sample_rate: Sample rate of the waveform. Required unless `audio`
|
|
28
|
+
carries one; given here it wins, for a track whose rate was
|
|
29
|
+
reported wrong
|
|
30
|
+
|
|
31
|
+
Returns:
|
|
32
|
+
One AudioVideo holding the frames and the track
|
|
33
|
+
|
|
34
|
+
Raises:
|
|
35
|
+
ValueError: If no waveform was given, or if no sample rate can be
|
|
36
|
+
established for the one that was
|
|
37
|
+
"""
|
|
38
|
+
waveform = getattr(audio, "audio", audio)
|
|
39
|
+
if waveform is None:
|
|
40
|
+
raise ValueError(
|
|
41
|
+
"pair_audio needs an audio track - the video it was given carries none"
|
|
42
|
+
)
|
|
43
|
+
|
|
44
|
+
rate = (
|
|
45
|
+
sample_rate if sample_rate is not None else getattr(audio, "sample_rate", None)
|
|
46
|
+
)
|
|
47
|
+
if rate is None:
|
|
48
|
+
raise ValueError(
|
|
49
|
+
"pair_audio needs 'sample_rate' - the audio it was given does not "
|
|
50
|
+
"carry one of its own"
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
# Frames are left in whatever shape they arrived in - the result saves a frame
|
|
54
|
+
# list, an array and a tensor alike, and converting a long video here would
|
|
55
|
+
# cost a copy of the whole thing for nothing
|
|
56
|
+
frames = video.frames if isinstance(video, AudioVideo) else video
|
|
57
|
+
logger.debug(f"Pairing frames with audio at {rate} Hz")
|
|
58
|
+
return AudioVideo(frames, as_channels_samples(waveform), rate)
|
dw/tasks/qr_code.py
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
import qrcode
|
|
2
|
+
from .image_utils import resize_resample
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
def get_qrcode_image(qr_code_contents, height=768, width=768):
|
|
6
|
+
# base the resolution off of size - defaulting to 768
|
|
7
|
+
resolution = max(height, width)
|
|
8
|
+
|
|
9
|
+
qr = qrcode.QRCode(
|
|
10
|
+
version=None,
|
|
11
|
+
error_correction=qrcode.constants.ERROR_CORRECT_H,
|
|
12
|
+
box_size=10,
|
|
13
|
+
border=4,
|
|
14
|
+
)
|
|
15
|
+
qr.add_data(qr_code_contents)
|
|
16
|
+
qr.make(fit=True)
|
|
17
|
+
|
|
18
|
+
qrcode_image = qr.make_image(fill_color="black", back_color="white")
|
|
19
|
+
return resize_resample(qrcode_image, resolution)
|
|
@@ -0,0 +1,175 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Face restoration via spandrel + facexlib.
|
|
3
|
+
|
|
4
|
+
Uses facexlib for face detection/alignment/pasting and spandrel for
|
|
5
|
+
neural network inference on cropped faces. Supports GFPGAN, RestoreFormer,
|
|
6
|
+
and CodeFormer (via spandrel-extra-arches) model weights.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import logging
|
|
10
|
+
import numpy as np
|
|
11
|
+
import torch
|
|
12
|
+
from PIL import Image
|
|
13
|
+
|
|
14
|
+
from .model_cache import cached_model
|
|
15
|
+
|
|
16
|
+
logger = logging.getLogger("dw")
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def restore_faces(image, model_name, device="cpu", **kwargs):
|
|
20
|
+
"""Restore faces in an image using a spandrel-compatible face restoration model.
|
|
21
|
+
|
|
22
|
+
Args:
|
|
23
|
+
image: PIL Image containing faces to restore
|
|
24
|
+
model_name: HuggingFace repo ID or local path to model weights
|
|
25
|
+
device: Target device ("cuda", "mps", "cpu")
|
|
26
|
+
**kwargs:
|
|
27
|
+
filename: Weight file name within a HF repo (default: auto-detect)
|
|
28
|
+
upscale_factor: Background upscale factor (default: 1, no upscaling)
|
|
29
|
+
face_size: Cropped face size in pixels (default: 512)
|
|
30
|
+
use_parse: Use face parsing for better blending (default: True)
|
|
31
|
+
only_center_face: Only restore the largest/center face (default: False)
|
|
32
|
+
detection_resize: Resize shorter side for detection speed (default: 640)
|
|
33
|
+
eye_dist_threshold: Skip faces with eye distance below this (default: 5)
|
|
34
|
+
upsample_img: Pre-upscaled background PIL Image (default: None)
|
|
35
|
+
|
|
36
|
+
Returns:
|
|
37
|
+
PIL Image with restored faces
|
|
38
|
+
"""
|
|
39
|
+
try:
|
|
40
|
+
from facexlib.utils.face_restoration_helper import FaceRestoreHelper
|
|
41
|
+
except ImportError:
|
|
42
|
+
raise ImportError(
|
|
43
|
+
"facexlib is required for face restoration. Install with: pip install facexlib"
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
from .upscale import _resolve_model_path
|
|
47
|
+
|
|
48
|
+
filename = kwargs.get("filename", None)
|
|
49
|
+
upscale_factor = kwargs.get("upscale_factor", 1)
|
|
50
|
+
face_size = kwargs.get("face_size", 512)
|
|
51
|
+
use_parse = kwargs.get("use_parse", True)
|
|
52
|
+
only_center_face = kwargs.get("only_center_face", False)
|
|
53
|
+
detection_resize = kwargs.get("detection_resize", 640)
|
|
54
|
+
eye_dist_threshold = kwargs.get("eye_dist_threshold", 5)
|
|
55
|
+
upsample_img = kwargs.get("upsample_img", None)
|
|
56
|
+
|
|
57
|
+
# Load the face restoration model via spandrel
|
|
58
|
+
def load_descriptor():
|
|
59
|
+
model_path = _resolve_model_path(model_name, filename)
|
|
60
|
+
result = _load_face_model(model_path, device)
|
|
61
|
+
if device != "cpu" and result.supports_half:
|
|
62
|
+
result.model.half()
|
|
63
|
+
return result
|
|
64
|
+
|
|
65
|
+
descriptor = cached_model(
|
|
66
|
+
("restore_faces", model_name, filename, str(device)), load_descriptor
|
|
67
|
+
)
|
|
68
|
+
|
|
69
|
+
# Set up facexlib helper
|
|
70
|
+
face_helper = FaceRestoreHelper(
|
|
71
|
+
upscale_factor=upscale_factor,
|
|
72
|
+
face_size=face_size,
|
|
73
|
+
crop_ratio=(1, 1),
|
|
74
|
+
det_model="retinaface_resnet50",
|
|
75
|
+
use_parse=use_parse,
|
|
76
|
+
device=torch.device(device),
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
# Convert PIL to BGR numpy (facexlib format)
|
|
80
|
+
input_bgr = np.array(image.convert("RGB"))[:, :, ::-1].copy()
|
|
81
|
+
face_helper.read_image(input_bgr)
|
|
82
|
+
|
|
83
|
+
# Detect faces
|
|
84
|
+
num_faces = face_helper.get_face_landmarks_5(
|
|
85
|
+
only_center_face=only_center_face,
|
|
86
|
+
resize=detection_resize,
|
|
87
|
+
eye_dist_threshold=eye_dist_threshold,
|
|
88
|
+
)
|
|
89
|
+
logger.info(f"Detected {num_faces} face(s)")
|
|
90
|
+
|
|
91
|
+
if num_faces == 0:
|
|
92
|
+
logger.warning("No faces detected, returning original image")
|
|
93
|
+
return image
|
|
94
|
+
|
|
95
|
+
# Align and warp faces to face_size x face_size
|
|
96
|
+
face_helper.align_warp_face()
|
|
97
|
+
|
|
98
|
+
# Half precision was applied at load time (see load_descriptor above) if supported
|
|
99
|
+
use_half = device != "cpu" and descriptor.supports_half
|
|
100
|
+
model_dtype = torch.float16 if use_half else torch.float32
|
|
101
|
+
|
|
102
|
+
# Restore each face
|
|
103
|
+
for i, cropped_face in enumerate(face_helper.cropped_faces):
|
|
104
|
+
logger.debug(f"Restoring face {i + 1}/{num_faces}")
|
|
105
|
+
|
|
106
|
+
# BGR uint8 numpy -> float32 tensor [1, 3, H, W]
|
|
107
|
+
face_tensor = (
|
|
108
|
+
torch.from_numpy(cropped_face.astype(np.float32) / 255.0)
|
|
109
|
+
.permute(2, 0, 1)
|
|
110
|
+
.unsqueeze(0)
|
|
111
|
+
)
|
|
112
|
+
face_tensor = face_tensor.to(device=device, dtype=model_dtype)
|
|
113
|
+
|
|
114
|
+
with torch.inference_mode():
|
|
115
|
+
restored_tensor = descriptor(face_tensor)
|
|
116
|
+
|
|
117
|
+
# Tensor -> BGR uint8 numpy. Rounds rather than truncates when
|
|
118
|
+
# quantizing (matches diffusers' VaeImageProcessor.numpy_to_pil
|
|
119
|
+
# behavior) so exact 8-bit values don't drift down a level; this
|
|
120
|
+
# stays local rather than using tensor_image's shared helpers
|
|
121
|
+
# because the data here is BGR numpy (facexlib's format), not PIL/RGB.
|
|
122
|
+
restored = restored_tensor.squeeze(0).permute(1, 2, 0)
|
|
123
|
+
restored = restored.mul(255).round().clamp(0, 255).byte().cpu().numpy()
|
|
124
|
+
|
|
125
|
+
# Resize to expected face_size if model output differs
|
|
126
|
+
if restored.shape[:2] != (face_size, face_size):
|
|
127
|
+
import cv2
|
|
128
|
+
|
|
129
|
+
restored = cv2.resize(
|
|
130
|
+
restored, (face_size, face_size), interpolation=cv2.INTER_LANCZOS4
|
|
131
|
+
)
|
|
132
|
+
|
|
133
|
+
face_helper.add_restored_face(restored)
|
|
134
|
+
|
|
135
|
+
# Prepare inverse affine transforms
|
|
136
|
+
face_helper.get_inverse_affine()
|
|
137
|
+
|
|
138
|
+
# Paste faces back onto the image
|
|
139
|
+
upsample_bgr = None
|
|
140
|
+
if upsample_img is not None:
|
|
141
|
+
upsample_bgr = np.array(upsample_img.convert("RGB"))[:, :, ::-1].copy()
|
|
142
|
+
|
|
143
|
+
result_bgr = face_helper.paste_faces_to_input_image(upsample_img=upsample_bgr)
|
|
144
|
+
|
|
145
|
+
# BGR numpy -> PIL RGB
|
|
146
|
+
result = Image.fromarray(result_bgr[:, :, ::-1])
|
|
147
|
+
logger.info(
|
|
148
|
+
f"Face restoration complete ({num_faces} face(s), {result.width}x{result.height})"
|
|
149
|
+
)
|
|
150
|
+
|
|
151
|
+
face_helper.clean_all()
|
|
152
|
+
return result
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
def _load_face_model(model_path, device):
|
|
156
|
+
"""Load a face restoration model via spandrel."""
|
|
157
|
+
try:
|
|
158
|
+
from spandrel import ModelLoader, ImageModelDescriptor
|
|
159
|
+
except ImportError:
|
|
160
|
+
raise ImportError(
|
|
161
|
+
"spandrel is required for face restoration. Install with: pip install spandrel"
|
|
162
|
+
)
|
|
163
|
+
|
|
164
|
+
logger.info(f"Loading face restoration model from {model_path}")
|
|
165
|
+
loader = ModelLoader(device=torch.device(device))
|
|
166
|
+
descriptor = loader.load_from_file(model_path)
|
|
167
|
+
|
|
168
|
+
if not isinstance(descriptor, ImageModelDescriptor):
|
|
169
|
+
raise ValueError(
|
|
170
|
+
f"Model is not an image model (got {type(descriptor).__name__}). "
|
|
171
|
+
f"Expected a face restoration model (GFPGAN, CodeFormer, RestoreFormer)."
|
|
172
|
+
)
|
|
173
|
+
|
|
174
|
+
logger.info(f"Loaded {descriptor.architecture.name}")
|
|
175
|
+
return descriptor
|