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,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
@@ -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