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,661 @@
1
+ from PIL import Image
2
+ import numpy as np
3
+ from .borders import add_border_and_mask, add_border_and_mask_with_size
4
+ from .model_cache import cached_model
5
+ import torch
6
+
7
+ # cv2, controlnet_aux, transformers and the model-backed task modules are imported
8
+ # inside the functions that use them - at module scope they add seconds to every
9
+ # startup (including the REPL worker spawn and dw.validate) for workflows that
10
+ # never touch an image-processing task
11
+
12
+
13
+ def _import_controlnet_aux():
14
+ import controlnet_aux
15
+
16
+ return controlnet_aux
17
+
18
+
19
+ # ---------------------------------------------------------------------------
20
+ # controlnet_aux detector dispatch
21
+ #
22
+ # Most controlnet_aux detectors follow one shape:
23
+ # getattr(controlnet_aux, attr).from_pretrained(*repo_args, **repo_kwargs).to(device)(image, **call_kwargs, **kwargs)
24
+ # The table below drives a single generic loader for that shape instead of a
25
+ # hand-written branch per detector. Loaded detectors are cached by
26
+ # (attr, device) via cached_model so repeated process_image calls - e.g. once
27
+ # per cartesian-product iteration in step.py - reuse the loaded weights
28
+ # instead of reloading them from disk every time.
29
+ # ---------------------------------------------------------------------------
30
+
31
+ # name -> (controlnet_aux attribute, from_pretrained positional args,
32
+ # from_pretrained kwargs, fixed call() kwargs)
33
+ _PRETRAINED_DETECTOR_SPECS = {
34
+ "mlsd": ("MLSDdetector", ("lllyasviel/Annotators",), {}, {}),
35
+ "normal_bae": ("NormalBaeDetector", ("lllyasviel/Annotators",), {}, {}),
36
+ "lineart": ("LineartDetector", ("lllyasviel/Annotators",), {}, {"coarse": True}),
37
+ "openpose": (
38
+ "OpenposeDetector",
39
+ ("lllyasviel/Annotators",),
40
+ {},
41
+ {"hand_and_face": True},
42
+ ),
43
+ "hed": ("HEDdetector", ("lllyasviel/Annotators",), {}, {"scribble": False}),
44
+ "scribble": ("HEDdetector", ("lllyasviel/Annotators",), {}, {"scribble": True}),
45
+ "pidi": ("PidiNetDetector", ("lllyasviel/Annotators",), {}, {"safe": True}),
46
+ "midas": ("MidasDetector", ("lllyasviel/Annotators",), {}, {}),
47
+ "zoe": ("ZoeDetector", ("lllyasviel/Annotators",), {}, {}),
48
+ "teed": ("TEEDdetector", ("fal-ai/teed",), {"filename": "5_model.pth"}, {}),
49
+ "anyline": (
50
+ "AnylineDetector",
51
+ ("TheMistoAI/MistoLine",),
52
+ {"filename": "MTEED.pth", "subfolder": "Anyline"},
53
+ {},
54
+ ),
55
+ "leres": ("LeresDetector", ("lllyasviel/Annotators",), {}, {}),
56
+ # sam is the one from_pretrained detector that is never moved to device -
57
+ # see _PROCESSORS registration below.
58
+ }
59
+
60
+ # Detectors constructed with no from_pretrained call at all.
61
+ _ZERO_ARG_DETECTOR_SPECS = {
62
+ "shuffle": "ContentShuffleDetector",
63
+ # controlnet_aux CannyDetector: plain cv2 Canny internally, but resizes
64
+ # the input to 512px first. Kept distinct from "canny_cv" below, which
65
+ # runs cv2 directly at the image's native resolution - same algorithm,
66
+ # different output size, so both names are intentional, not duplicates.
67
+ "canny": "CannyDetector",
68
+ "lineart_standard": "LineartStandardDetector",
69
+ }
70
+
71
+
72
+ def _build_pretrained_detector(attr, repo_args, repo_kwargs, to_device, device):
73
+ controlnet_aux = _import_controlnet_aux()
74
+ detector = getattr(controlnet_aux, attr).from_pretrained(*repo_args, **repo_kwargs)
75
+ if to_device:
76
+ detector = detector.to(device)
77
+ return detector
78
+
79
+
80
+ def _build_zero_arg_detector(attr):
81
+ controlnet_aux = _import_controlnet_aux()
82
+ return getattr(controlnet_aux, attr)()
83
+
84
+
85
+ def _make_pretrained_handler(attr, repo_args, repo_kwargs, call_kwargs, to_device=True):
86
+ def handler(image, device, kwargs):
87
+ detector = cached_model(
88
+ ("image_processor", attr, str(device)),
89
+ lambda: _build_pretrained_detector(
90
+ attr, repo_args, repo_kwargs, to_device, device
91
+ ),
92
+ )
93
+ return detector(image, **call_kwargs, **kwargs)
94
+
95
+ return handler
96
+
97
+
98
+ def _make_zero_arg_handler(attr):
99
+ def handler(image, device, kwargs):
100
+ detector = cached_model(
101
+ ("image_processor", attr, str(device)),
102
+ lambda: _build_zero_arg_detector(attr),
103
+ )
104
+ return detector(image, **kwargs)
105
+
106
+ return handler
107
+
108
+
109
+ def _dw_pose_handler(image, device, kwargs):
110
+ detector = cached_model(
111
+ ("image_processor", "DWposeDetector", str(device)),
112
+ lambda: _import_controlnet_aux().DWposeDetector(device=device),
113
+ )
114
+ return detector(image, **kwargs)
115
+
116
+
117
+ def _remove_background_handler(image, device, kwargs):
118
+ from .background_remover import remove_background
119
+
120
+ return remove_background(image, device, **kwargs)
121
+
122
+
123
+ def _depth_estimator_tensor_handler(image, device, kwargs):
124
+ from .depth_estimator import make_hint_tensor
125
+
126
+ return make_hint_tensor(image, device, **kwargs)
127
+
128
+
129
+ def _depth_estimator_handler(image, device, kwargs):
130
+ from .depth_estimator import make_hint_image
131
+
132
+ return make_hint_image(image, device, **kwargs)
133
+
134
+
135
+ def get_zoe_depth_map(image, device):
136
+ from .zoe_depth import colorize, load_zoe
137
+
138
+ model_zoe_n = load_zoe(device)
139
+ # MPS doesn't support autocast, so use 'cpu' for autocast when on MPS
140
+ from dw import get_autocast_device_type
141
+
142
+ autocast_device = get_autocast_device_type()
143
+ if autocast_device == "cuda":
144
+ with torch.autocast(autocast_device, enabled=True):
145
+ depth = model_zoe_n.infer_pil(image)
146
+ else:
147
+ # For MPS/CPU, don't use autocast
148
+ depth = model_zoe_n.infer_pil(image)
149
+ return colorize(depth, cmap="gray_r")
150
+
151
+
152
+ def image_to_canny(image, low_threshold=100, high_threshold=200):
153
+ # Raw cv2.Canny at the image's native resolution - intentionally kept
154
+ # separate from the "canny" controlnet_aux CannyDetector above, which
155
+ # resizes to 512px first. See comment on _ZERO_ARG_DETECTOR_SPECS["canny"].
156
+ import cv2
157
+
158
+ image = np.array(image)
159
+
160
+ image = cv2.Canny(image, low_threshold, high_threshold)
161
+ image = image[:, :, None]
162
+ image = np.concatenate([image, image, image], axis=2)
163
+ return Image.fromarray(image)
164
+
165
+
166
+ def image_to_depth(image, device, height=1024, width=1024):
167
+ from transformers import DPTForDepthEstimation, DPTImageProcessor
168
+
169
+ size = (width, height)
170
+ depth_estimator = DPTForDepthEstimation.from_pretrained(
171
+ "Intel/dpt-hybrid-midas"
172
+ ).to(device)
173
+ feature_extractor = DPTImageProcessor.from_pretrained("Intel/dpt-hybrid-midas")
174
+
175
+ image = feature_extractor(images=image, return_tensors="pt").pixel_values.to(device)
176
+ # MPS doesn't support autocast, so use 'cpu' for autocast when on MPS
177
+ from dw import get_autocast_device_type
178
+
179
+ autocast_device = get_autocast_device_type()
180
+ if autocast_device == "cuda":
181
+ with torch.no_grad(), torch.autocast(autocast_device):
182
+ depth_map = depth_estimator(image).predicted_depth
183
+ else:
184
+ # For MPS/CPU, don't use autocast
185
+ with torch.no_grad():
186
+ depth_map = depth_estimator(image).predicted_depth
187
+
188
+ depth_map = torch.nn.functional.interpolate(
189
+ depth_map.unsqueeze(1),
190
+ size=size,
191
+ mode="bicubic",
192
+ align_corners=False,
193
+ )
194
+ depth_min = torch.amin(depth_map, dim=[1, 2, 3], keepdim=True)
195
+ depth_max = torch.amax(depth_map, dim=[1, 2, 3], keepdim=True)
196
+ depth_map = (depth_map - depth_min) / (depth_max - depth_min)
197
+ image = torch.cat([depth_map] * 3, dim=1)
198
+
199
+ image = image.permute(0, 2, 3, 1).cpu().numpy()[0]
200
+ image = Image.fromarray((image * 255.0).clip(0, 255).astype(np.uint8))
201
+ return image
202
+
203
+
204
+ def image_to_segmentation(image):
205
+ from transformers import AutoImageProcessor, UperNetForSemanticSegmentation
206
+
207
+ image_processor = AutoImageProcessor.from_pretrained(
208
+ "openmmlab/upernet-convnext-small"
209
+ )
210
+ image_segmentor = UperNetForSemanticSegmentation.from_pretrained(
211
+ "openmmlab/upernet-convnext-small"
212
+ )
213
+ pixel_values = image_processor(image, return_tensors="pt").pixel_values
214
+ with torch.no_grad():
215
+ outputs = image_segmentor(pixel_values)
216
+ seg = image_processor.post_process_semantic_segmentation(
217
+ outputs, target_sizes=[image.size[::-1]]
218
+ )[0]
219
+ color_seg = np.zeros(
220
+ (seg.shape[0], seg.shape[1], 3), dtype=np.uint8
221
+ ) # height, width, 3
222
+ for label, color in enumerate(ada_palette):
223
+ color_seg[seg == label, :] = color
224
+ color_seg = color_seg.astype(np.uint8)
225
+ return Image.fromarray(color_seg)
226
+
227
+
228
+ def get_image_size(image):
229
+ return {"width": image.width, "height": image.height}
230
+
231
+
232
+ def crop_square(img: Image) -> Image:
233
+ # Determine the shortest side
234
+ min_side = min(img.width, img.height)
235
+
236
+ # Calculate the left and right crop positions for centering
237
+ left = (img.width - min_side) // 2
238
+ right = left + min_side
239
+
240
+ # Calculate the top and bottom crop positions for centering
241
+ top = (img.height - min_side) // 2
242
+ bottom = top + min_side
243
+
244
+ # Crop the image
245
+ img_cropped = img.crop((left, top, right, bottom))
246
+
247
+ return img_cropped
248
+
249
+
250
+ def resize_center_crop(img, height=768, width=768):
251
+ output_size = (width, height)
252
+ W, H = img.size
253
+
254
+ # Calculate dimensions to crop to the center
255
+ new_dimension = min(W, H)
256
+ left = (W - new_dimension) / 2
257
+ top = (H - new_dimension) / 2
258
+ right = (W + new_dimension) / 2
259
+ bottom = (H + new_dimension) / 2
260
+
261
+ # Crop and resize
262
+ img = img.crop((left, top, right, bottom))
263
+ img = img.resize(output_size)
264
+
265
+ return img
266
+
267
+
268
+ def resize_rescale(image, height=768, width=768):
269
+ input_image = image.convert("RGB")
270
+ return input_image.resize((width, height))
271
+
272
+
273
+ def resize_resample(image, resolution=1024):
274
+ input_image = image.convert("RGB")
275
+ W, H = input_image.size
276
+ k = float(resolution) / min(H, W)
277
+ H *= k
278
+ W *= k
279
+ H = int(round(H / 64.0)) * 64
280
+ W = int(round(W / 64.0)) * 64
281
+
282
+ return input_image.resize((W, H), resample=Image.LANCZOS)
283
+
284
+
285
+ # Standard aspect ratios used by SDXL, Flux, and similar models.
286
+ # Each entry is (width_ratio, height_ratio).
287
+ _DEFAULT_RATIOS = [
288
+ (1, 1),
289
+ (4, 3),
290
+ (3, 4),
291
+ (3, 2),
292
+ (2, 3),
293
+ (16, 9),
294
+ (9, 16),
295
+ (21, 9),
296
+ (9, 21),
297
+ ]
298
+
299
+
300
+ def resize_bucket(image, resolution=1024, ratios=None, alignment=64):
301
+ """Resize image to the closest model-native aspect ratio bucket.
302
+
303
+ Picks the standard ratio closest to the input image's natural aspect
304
+ ratio, then scales to fit within the target resolution (based on the
305
+ short side) with dimensions aligned to `alignment` pixels.
306
+
307
+ Args:
308
+ image: PIL Image to resize.
309
+ resolution: Target size for the short side in pixels (default: 1024).
310
+ ratios: Optional list of [w, h] ratio pairs. Defaults to standard
311
+ ratios used by SDXL/Flux (1:1, 4:3, 3:2, 16:9, etc.).
312
+ alignment: Round dimensions to this multiple (default: 64).
313
+
314
+ Returns:
315
+ PIL Image resized to the bucketed dimensions.
316
+ """
317
+ input_image = image.convert("RGB")
318
+ W, H = input_image.size
319
+ input_ratio = W / H
320
+
321
+ bucket_ratios = ratios if ratios is not None else _DEFAULT_RATIOS
322
+
323
+ # Find the closest aspect ratio
324
+ best_ratio = min(
325
+ bucket_ratios,
326
+ key=lambda r: abs((r[0] / r[1]) - input_ratio),
327
+ )
328
+
329
+ wr, hr = best_ratio
330
+ bucket_ratio = wr / hr
331
+
332
+ # Scale so the short side matches resolution, then align
333
+ if bucket_ratio >= 1.0:
334
+ # Landscape or square: height is the short side
335
+ out_h = int(round(resolution / alignment)) * alignment
336
+ out_w = int(round((out_h * bucket_ratio) / alignment)) * alignment
337
+ else:
338
+ # Portrait: width is the short side
339
+ out_w = int(round(resolution / alignment)) * alignment
340
+ out_h = int(round((out_w / bucket_ratio) / alignment)) * alignment
341
+
342
+ return input_image.resize((out_w, out_h), resample=Image.LANCZOS)
343
+
344
+
345
+ def strip_exif(image):
346
+ """Remove all EXIF and metadata from an image.
347
+
348
+ Creates a clean copy with pixel data only — no GPS coordinates,
349
+ camera info, timestamps, or other embedded metadata.
350
+
351
+ Args:
352
+ image: PIL Image to strip.
353
+
354
+ Returns:
355
+ PIL Image with all metadata removed.
356
+ """
357
+ clean = Image.new(image.mode, image.size)
358
+ clean.paste(image)
359
+ return clean
360
+
361
+
362
+ def add_watermark(
363
+ image,
364
+ text="AI Generated",
365
+ position="bottom-right",
366
+ opacity=128,
367
+ font_size=0,
368
+ margin=10,
369
+ color=None,
370
+ ):
371
+ """Add a visible text watermark to an image.
372
+
373
+ Args:
374
+ image: PIL Image to watermark.
375
+ text: Watermark text (default: "AI Generated").
376
+ position: Placement — "bottom-right", "bottom-left", "top-right",
377
+ "top-left", or "center" (default: "bottom-right").
378
+ opacity: Text opacity 0-255 (default: 128).
379
+ font_size: Font size in pixels. 0 = auto-scale to ~3% of image height.
380
+ margin: Pixel margin from edges (default: 10).
381
+ color: RGB tuple for text color (default: white).
382
+
383
+ Returns:
384
+ PIL Image with watermark applied.
385
+ """
386
+ from PIL import ImageDraw, ImageFont
387
+
388
+ base = image.convert("RGBA")
389
+ overlay = Image.new("RGBA", base.size, (0, 0, 0, 0))
390
+ draw = ImageDraw.Draw(overlay)
391
+
392
+ if color is None:
393
+ color = (255, 255, 255)
394
+ fill = (*color, int(opacity))
395
+
396
+ if font_size <= 0:
397
+ font_size = max(12, base.height // 30)
398
+
399
+ try:
400
+ font = ImageFont.truetype("Arial", font_size)
401
+ except (IOError, OSError):
402
+ font = ImageFont.load_default(size=font_size)
403
+
404
+ bbox = draw.textbbox((0, 0), text, font=font)
405
+ text_w = bbox[2] - bbox[0]
406
+ text_h = bbox[3] - bbox[1]
407
+
408
+ positions = {
409
+ "bottom-right": (base.width - text_w - margin, base.height - text_h - margin),
410
+ "bottom-left": (margin, base.height - text_h - margin),
411
+ "top-right": (base.width - text_w - margin, margin),
412
+ "top-left": (margin, margin),
413
+ "center": ((base.width - text_w) // 2, (base.height - text_h) // 2),
414
+ }
415
+ xy = positions.get(position, positions["bottom-right"])
416
+
417
+ draw.text(xy, text, font=font, fill=fill)
418
+
419
+ result = Image.alpha_composite(base, overlay)
420
+ return result.convert("RGB")
421
+
422
+
423
+ # ---------------------------------------------------------------------------
424
+ # process_image dispatch table
425
+ #
426
+ # Every handler has the uniform signature (image, device, kwargs) -> result,
427
+ # so process_image is just a lookup + call. Built once at import time from
428
+ # the detector spec tables above plus direct entries for the plain PIL/task
429
+ # functions.
430
+ # ---------------------------------------------------------------------------
431
+
432
+ _PROCESSORS = {
433
+ "get_image_size": lambda image, device, kwargs: get_image_size(image),
434
+ "add_border_and_mask": lambda image, device, kwargs: add_border_and_mask(
435
+ image, **kwargs
436
+ ),
437
+ "add_border_and_mask_with_size": lambda image, device, kwargs: add_border_and_mask_with_size(
438
+ image, **kwargs
439
+ ),
440
+ "remove_background": _remove_background_handler,
441
+ # Raw cv2 Canny at native resolution - see image_to_canny() docstring
442
+ # comment for how this differs from "canny" below.
443
+ "canny_cv": lambda image, device, kwargs: image_to_canny(image, **kwargs),
444
+ "segmentation": lambda image, device, kwargs: image_to_segmentation(image),
445
+ "zoe_depth": lambda image, device, kwargs: get_zoe_depth_map(image, device),
446
+ "depth": lambda image, device, kwargs: image_to_depth(image, device, **kwargs),
447
+ "depth_estimator_tensor": _depth_estimator_tensor_handler,
448
+ "depth_estimator": _depth_estimator_handler,
449
+ "resize_center_crop": lambda image, device, kwargs: resize_center_crop(
450
+ image, **kwargs
451
+ ),
452
+ "resize_resample": lambda image, device, kwargs: resize_resample(image, **kwargs),
453
+ "crop_square": lambda image, device, kwargs: crop_square(image, **kwargs),
454
+ "resize_rescale": lambda image, device, kwargs: resize_rescale(image, **kwargs),
455
+ "resize_bucket": lambda image, device, kwargs: resize_bucket(image, **kwargs),
456
+ "strip_exif": lambda image, device, kwargs: strip_exif(image),
457
+ "add_watermark": lambda image, device, kwargs: add_watermark(image, **kwargs),
458
+ }
459
+
460
+ for _name, (
461
+ _attr,
462
+ _repo_args,
463
+ _repo_kwargs,
464
+ _call_kwargs,
465
+ ) in _PRETRAINED_DETECTOR_SPECS.items():
466
+ _PROCESSORS[_name] = _make_pretrained_handler(
467
+ _attr, _repo_args, _repo_kwargs, _call_kwargs
468
+ )
469
+
470
+ # sam is the one from_pretrained detector never moved to device - matches
471
+ # the pre-refactor behavior, which called it straight off from_pretrained().
472
+ _PROCESSORS["sam"] = _make_pretrained_handler(
473
+ "SamDetector",
474
+ ("ybelkada/segment-anything",),
475
+ {"subfolder": "checkpoints"},
476
+ {},
477
+ to_device=False,
478
+ )
479
+
480
+ for _name, _attr in _ZERO_ARG_DETECTOR_SPECS.items():
481
+ _PROCESSORS[_name] = _make_zero_arg_handler(_attr)
482
+
483
+ _PROCESSORS["dw_pose"] = _dw_pose_handler
484
+
485
+ del _name, _attr, _repo_args, _repo_kwargs, _call_kwargs
486
+
487
+
488
+ def available_processors():
489
+ """Return the sorted list of processor names process_image accepts.
490
+
491
+ Used by command registration to enumerate supported image processors
492
+ without duplicating this dispatch table.
493
+ """
494
+ return sorted(_PROCESSORS)
495
+
496
+
497
+ def process_image(image, processor, device, kwargs):
498
+ processor = processor.lower()
499
+
500
+ handler = _PROCESSORS.get(processor)
501
+ if handler is None:
502
+ raise Exception(f"Unknown image processor type: {processor}")
503
+
504
+ return handler(image, device, kwargs)
505
+
506
+
507
+ ada_palette = np.asarray(
508
+ [
509
+ [0, 0, 0],
510
+ [120, 120, 120],
511
+ [180, 120, 120],
512
+ [6, 230, 230],
513
+ [80, 50, 50],
514
+ [4, 200, 3],
515
+ [120, 120, 80],
516
+ [140, 140, 140],
517
+ [204, 5, 255],
518
+ [230, 230, 230],
519
+ [4, 250, 7],
520
+ [224, 5, 255],
521
+ [235, 255, 7],
522
+ [150, 5, 61],
523
+ [120, 120, 70],
524
+ [8, 255, 51],
525
+ [255, 6, 82],
526
+ [143, 255, 140],
527
+ [204, 255, 4],
528
+ [255, 51, 7],
529
+ [204, 70, 3],
530
+ [0, 102, 200],
531
+ [61, 230, 250],
532
+ [255, 6, 51],
533
+ [11, 102, 255],
534
+ [255, 7, 71],
535
+ [255, 9, 224],
536
+ [9, 7, 230],
537
+ [220, 220, 220],
538
+ [255, 9, 92],
539
+ [112, 9, 255],
540
+ [8, 255, 214],
541
+ [7, 255, 224],
542
+ [255, 184, 6],
543
+ [10, 255, 71],
544
+ [255, 41, 10],
545
+ [7, 255, 255],
546
+ [224, 255, 8],
547
+ [102, 8, 255],
548
+ [255, 61, 6],
549
+ [255, 194, 7],
550
+ [255, 122, 8],
551
+ [0, 255, 20],
552
+ [255, 8, 41],
553
+ [255, 5, 153],
554
+ [6, 51, 255],
555
+ [235, 12, 255],
556
+ [160, 150, 20],
557
+ [0, 163, 255],
558
+ [140, 140, 140],
559
+ [250, 10, 15],
560
+ [20, 255, 0],
561
+ [31, 255, 0],
562
+ [255, 31, 0],
563
+ [255, 224, 0],
564
+ [153, 255, 0],
565
+ [0, 0, 255],
566
+ [255, 71, 0],
567
+ [0, 235, 255],
568
+ [0, 173, 255],
569
+ [31, 0, 255],
570
+ [11, 200, 200],
571
+ [255, 82, 0],
572
+ [0, 255, 245],
573
+ [0, 61, 255],
574
+ [0, 255, 112],
575
+ [0, 255, 133],
576
+ [255, 0, 0],
577
+ [255, 163, 0],
578
+ [255, 102, 0],
579
+ [194, 255, 0],
580
+ [0, 143, 255],
581
+ [51, 255, 0],
582
+ [0, 82, 255],
583
+ [0, 255, 41],
584
+ [0, 255, 173],
585
+ [10, 0, 255],
586
+ [173, 255, 0],
587
+ [0, 255, 153],
588
+ [255, 92, 0],
589
+ [255, 0, 255],
590
+ [255, 0, 245],
591
+ [255, 0, 102],
592
+ [255, 173, 0],
593
+ [255, 0, 20],
594
+ [255, 184, 184],
595
+ [0, 31, 255],
596
+ [0, 255, 61],
597
+ [0, 71, 255],
598
+ [255, 0, 204],
599
+ [0, 255, 194],
600
+ [0, 255, 82],
601
+ [0, 10, 255],
602
+ [0, 112, 255],
603
+ [51, 0, 255],
604
+ [0, 194, 255],
605
+ [0, 122, 255],
606
+ [0, 255, 163],
607
+ [255, 153, 0],
608
+ [0, 255, 10],
609
+ [255, 112, 0],
610
+ [143, 255, 0],
611
+ [82, 0, 255],
612
+ [163, 255, 0],
613
+ [255, 235, 0],
614
+ [8, 184, 170],
615
+ [133, 0, 255],
616
+ [0, 255, 92],
617
+ [184, 0, 255],
618
+ [255, 0, 31],
619
+ [0, 184, 255],
620
+ [0, 214, 255],
621
+ [255, 0, 112],
622
+ [92, 255, 0],
623
+ [0, 224, 255],
624
+ [112, 224, 255],
625
+ [70, 184, 160],
626
+ [163, 0, 255],
627
+ [153, 0, 255],
628
+ [71, 255, 0],
629
+ [255, 0, 163],
630
+ [255, 204, 0],
631
+ [255, 0, 143],
632
+ [0, 255, 235],
633
+ [133, 255, 0],
634
+ [255, 0, 235],
635
+ [245, 0, 255],
636
+ [255, 0, 122],
637
+ [255, 245, 0],
638
+ [10, 190, 212],
639
+ [214, 255, 0],
640
+ [0, 204, 255],
641
+ [20, 0, 255],
642
+ [255, 255, 0],
643
+ [0, 153, 255],
644
+ [0, 41, 255],
645
+ [0, 255, 204],
646
+ [41, 0, 255],
647
+ [41, 255, 0],
648
+ [173, 0, 255],
649
+ [0, 245, 255],
650
+ [71, 0, 255],
651
+ [122, 0, 255],
652
+ [0, 255, 184],
653
+ [0, 92, 255],
654
+ [184, 255, 0],
655
+ [0, 133, 255],
656
+ [255, 214, 0],
657
+ [25, 194, 194],
658
+ [102, 255, 0],
659
+ [92, 0, 255],
660
+ ]
661
+ )