diffusers-workflow 0.4.0__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 (260) hide show
  1. diffusers_workflow-0.4.0.dist-info/METADATA +318 -0
  2. diffusers_workflow-0.4.0.dist-info/RECORD +260 -0
  3. diffusers_workflow-0.4.0.dist-info/WHEEL +5 -0
  4. diffusers_workflow-0.4.0.dist-info/entry_points.txt +7 -0
  5. diffusers_workflow-0.4.0.dist-info/licenses/LICENSE +201 -0
  6. diffusers_workflow-0.4.0.dist-info/top_level.txt +2 -0
  7. dw/__init__.py +440 -0
  8. dw/adapter_compatibility.py +226 -0
  9. dw/arguments.py +1231 -0
  10. dw/assessment_rules.py +159 -0
  11. dw/assets.py +130 -0
  12. dw/cache_blocks.json +16 -0
  13. dw/cache_blocks.py +146 -0
  14. dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
  15. dw/content_types.py +150 -0
  16. dw/dissolve_frame_errors.py +121 -0
  17. dw/docs/ACCELERATION.md +352 -0
  18. dw/docs/AGENT_LOOP.md +95 -0
  19. dw/docs/DEPENDENCIES.md +91 -0
  20. dw/docs/IP_ADAPTER.md +109 -0
  21. dw/docs/LORAS.md +131 -0
  22. dw/docs/MCP.md +517 -0
  23. dw/docs/PROMPT_WEIGHTING.md +78 -0
  24. dw/docs/QUANTIZATION.md +230 -0
  25. dw/docs/RECIPES_24GB.md +201 -0
  26. dw/docs/RELEASING.md +195 -0
  27. dw/docs/REMOTE.md +140 -0
  28. dw/docs/REPL_COMMANDS.md +121 -0
  29. dw/docs/REPL_WORKER_GUIDE.md +51 -0
  30. dw/docs/SECURITY.md +272 -0
  31. dw/docs/SECURITY_QUICKREF.md +112 -0
  32. dw/docs/SERVER.md +679 -0
  33. dw/docs/TASKS.md +1741 -0
  34. dw/docs/TESTING.md +71 -0
  35. dw/docs/WORKFLOW_GUIDE.md +2038 -0
  36. dw/docs/WORKSPACES.md +316 -0
  37. dw/download_watch.py +335 -0
  38. dw/elision.py +306 -0
  39. dw/events.py +275 -0
  40. dw/for_each.py +409 -0
  41. dw/host_memory.py +258 -0
  42. dw/host_memory_projection.py +230 -0
  43. dw/hub_cache.py +432 -0
  44. dw/introspection.py +1228 -0
  45. dw/kernel_availability.py +208 -0
  46. dw/locations.py +599 -0
  47. dw/log_setup.py +45 -0
  48. dw/loudness.py +82 -0
  49. dw/media_audio.py +217 -0
  50. dw/media_frames.py +367 -0
  51. dw/media_info.py +297 -0
  52. dw/pipeline_processors/chain.py +821 -0
  53. dw/pipeline_processors/config_objects.py +237 -0
  54. dw/pipeline_processors/pipeline.py +2297 -0
  55. dw/pipeline_processors/remote.py +46 -0
  56. dw/plan.py +920 -0
  57. dw/previous_results.py +411 -0
  58. dw/probe_paths.py +59 -0
  59. dw/prompt_schema.json +48 -0
  60. dw/prompt_weighting.py +378 -0
  61. dw/prompts.py +159 -0
  62. dw/realize.py +250 -0
  63. dw/reference_limits.py +215 -0
  64. dw/reference_names.py +125 -0
  65. dw/repl.py +338 -0
  66. dw/repl_commands.py +836 -0
  67. dw/repl_worker.py +159 -0
  68. dw/result.py +1720 -0
  69. dw/result_fps.py +82 -0
  70. dw/run.py +162 -0
  71. dw/runs.py +768 -0
  72. dw/scalar_result_validation.py +97 -0
  73. dw/schema.py +283 -0
  74. dw/security.py +1038 -0
  75. dw/select_validation.py +115 -0
  76. dw/serve.py +277 -0
  77. dw/server/__init__.py +2 -0
  78. dw/server/app.py +4586 -0
  79. dw/server/assess.py +132 -0
  80. dw/server/catalog_shape.py +487 -0
  81. dw/server/enhancers.py +129 -0
  82. dw/server/exports.py +480 -0
  83. dw/server/guides.py +257 -0
  84. dw/server/jobs.py +1561 -0
  85. dw/server/mcp_mount.py +95 -0
  86. dw/server/netinfo.py +124 -0
  87. dw/server/observed_cost.py +379 -0
  88. dw/server/sysinfo.py +71 -0
  89. dw/server/ui/assets/abap-08VXUWAP.js +1 -0
  90. dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
  91. dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
  92. dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
  93. dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
  94. dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
  95. dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
  96. dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
  97. dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
  98. dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
  99. dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
  100. dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
  101. dw/server/ui/assets/css-DIMkf-bt.js +3 -0
  102. dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
  103. dw/server/ui/assets/cssMode-CPznxfY8.js +1 -0
  104. dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
  105. dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
  106. dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
  107. dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
  108. dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
  109. dw/server/ui/assets/editor.api-CpWcotrd.js +847 -0
  110. dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
  111. dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
  112. dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
  113. dw/server/ui/assets/freemarker2-CXtRM8N4.js +3 -0
  114. dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
  115. dw/server/ui/assets/go-C-y9NEjX.js +1 -0
  116. dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
  117. dw/server/ui/assets/handlebars-N7x-6NMY.js +1 -0
  118. dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
  119. dw/server/ui/assets/html-PhsdjHSr.js +1 -0
  120. dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
  121. dw/server/ui/assets/htmlMode-Dgj0SEok.js +1 -0
  122. dw/server/ui/assets/index-3Vw6WAPW.css +1 -0
  123. dw/server/ui/assets/index-DgrYhQd9.js +43 -0
  124. dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
  125. dw/server/ui/assets/java-BEtHBSE6.js +1 -0
  126. dw/server/ui/assets/javascript-BJqN9Qhv.js +1 -0
  127. dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
  128. dw/server/ui/assets/jsonMode-DbM4SWSv.js +7 -0
  129. dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
  130. dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
  131. dw/server/ui/assets/less-B9JPFI3C.js +2 -0
  132. dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
  133. dw/server/ui/assets/liquid-BWr8lEc4.js +1 -0
  134. dw/server/ui/assets/lspLanguageFeatures-C1iGuDyZ.js +4 -0
  135. dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
  136. dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
  137. dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
  138. dw/server/ui/assets/mdx-DAdMi_0p.js +1 -0
  139. dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
  140. dw/server/ui/assets/monaco--ixms01u.css +1 -0
  141. dw/server/ui/assets/monaco-BGCeEqaw.js +56 -0
  142. dw/server/ui/assets/msdax-DauUninz.js +1 -0
  143. dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
  144. dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
  145. dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
  146. dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
  147. dw/server/ui/assets/perl-oz_6vUea.js +1 -0
  148. dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
  149. dw/server/ui/assets/php-nr791fC2.js +1 -0
  150. dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
  151. dw/server/ui/assets/postiats-43DmfD33.js +1 -0
  152. dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
  153. dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
  154. dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
  155. dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
  156. dw/server/ui/assets/python-Bcn70HdC.js +1 -0
  157. dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
  158. dw/server/ui/assets/r-BwWrilGY.js +1 -0
  159. dw/server/ui/assets/razor-D1HmNnby.js +1 -0
  160. dw/server/ui/assets/redis-ClamHrr6.js +1 -0
  161. dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
  162. dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
  163. dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
  164. dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
  165. dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
  166. dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
  167. dw/server/ui/assets/scheme-BeGwcela.js +1 -0
  168. dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
  169. dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
  170. dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
  171. dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
  172. dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
  173. dw/server/ui/assets/sql-NEE52Syq.js +1 -0
  174. dw/server/ui/assets/st-DbInun42.js +1 -0
  175. dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
  176. dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
  177. dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
  178. dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
  179. dw/server/ui/assets/tsMode-D6u0XmOW.js +11 -0
  180. dw/server/ui/assets/twig-De2hgUGE.js +1 -0
  181. dw/server/ui/assets/typescript-BU6v-LMV.js +1 -0
  182. dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
  183. dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
  184. dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
  185. dw/server/ui/assets/workers-Cn7cTUKr.js +1 -0
  186. dw/server/ui/assets/xml--0LP2Lwk.js +1 -0
  187. dw/server/ui/assets/yaml-mpBg9jnt.js +1 -0
  188. dw/server/ui/index.html +17 -0
  189. dw/server/updater.py +192 -0
  190. dw/settings.py +98 -0
  191. dw/shot_span_preflight.py +116 -0
  192. dw/shots.py +359 -0
  193. dw/slice_preflight.py +148 -0
  194. dw/step.py +187 -0
  195. dw/step_cache.py +442 -0
  196. dw/subfolders.py +107 -0
  197. dw/task_domains.py +307 -0
  198. dw/tasks/assess.py +826 -0
  199. dw/tasks/audio_transcription.py +88 -0
  200. dw/tasks/audio_utils.py +1862 -0
  201. dw/tasks/background_remover.py +43 -0
  202. dw/tasks/borders.py +113 -0
  203. dw/tasks/compose_text.py +74 -0
  204. dw/tasks/concat_videos.py +300 -0
  205. dw/tasks/depth_estimator.py +54 -0
  206. dw/tasks/diffusion_upscale.py +109 -0
  207. dw/tasks/dissolve_videos.py +342 -0
  208. dw/tasks/format_messages.py +24 -0
  209. dw/tasks/gather.py +173 -0
  210. dw/tasks/grade.py +97 -0
  211. dw/tasks/image_to_text.py +43 -0
  212. dw/tasks/image_utils.py +764 -0
  213. dw/tasks/interpolate_frames.py +252 -0
  214. dw/tasks/judge.py +68 -0
  215. dw/tasks/model_cache.py +55 -0
  216. dw/tasks/pair_audio.py +268 -0
  217. dw/tasks/qr_code.py +19 -0
  218. dw/tasks/restore_faces.py +175 -0
  219. dw/tasks/rife_model.py +192 -0
  220. dw/tasks/segment.py +121 -0
  221. dw/tasks/select.py +111 -0
  222. dw/tasks/speech_generation.py +228 -0
  223. dw/tasks/stabilize.py +129 -0
  224. dw/tasks/task.py +920 -0
  225. dw/tasks/tensor_image.py +57 -0
  226. dw/tasks/text_generation.py +169 -0
  227. dw/tasks/text_sections.py +80 -0
  228. dw/tasks/upscale.py +203 -0
  229. dw/tasks/video_utils.py +624 -0
  230. dw/tasks/zoe_depth.py +71 -0
  231. dw/teacache.py +381 -0
  232. dw/teacache_models.json +99 -0
  233. dw/test.py +29 -0
  234. dw/type_helpers.py +231 -0
  235. dw/validate.py +68 -0
  236. dw/variable_constraints.py +444 -0
  237. dw/variables.py +443 -0
  238. dw/video_extensions.py +141 -0
  239. dw/vram_estimate.py +116 -0
  240. dw/worker.py +764 -0
  241. dw/workflow.py +2007 -0
  242. dw/workflow_schema.json +1346 -0
  243. dw/workflow_sources.py +383 -0
  244. dw/workflows/h3_context_ir.json +57 -0
  245. dw/workflows/test.json +31 -0
  246. dw/workspace.py +730 -0
  247. dw_mcp/__init__.py +6 -0
  248. dw_mcp/__main__.py +133 -0
  249. dw_mcp/assets.py +336 -0
  250. dw_mcp/authoring.py +114 -0
  251. dw_mcp/catalog.py +360 -0
  252. dw_mcp/client.py +486 -0
  253. dw_mcp/diagnose.py +371 -0
  254. dw_mcp/exports.py +84 -0
  255. dw_mcp/guides.py +35 -0
  256. dw_mcp/media.py +638 -0
  257. dw_mcp/models.py +97 -0
  258. dw_mcp/prompts.py +104 -0
  259. dw_mcp/server.py +1343 -0
  260. dw_mcp/workspaces.py +212 -0
@@ -0,0 +1,764 @@
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 edge map at the image's native resolution."""
154
+ # Raw cv2.Canny at the image's native resolution - intentionally kept
155
+ # separate from the "canny" controlnet_aux CannyDetector above, which
156
+ # resizes to 512px first. See comment on _ZERO_ARG_DETECTOR_SPECS["canny"].
157
+ import cv2
158
+
159
+ image = np.array(image)
160
+
161
+ image = cv2.Canny(image, low_threshold, high_threshold)
162
+ image = image[:, :, None]
163
+ image = np.concatenate([image, image, image], axis=2)
164
+ return Image.fromarray(image)
165
+
166
+
167
+ def image_to_depth(image, device, height=1024, width=1024):
168
+ from transformers import DPTForDepthEstimation, DPTImageProcessor
169
+
170
+ size = (width, height)
171
+ depth_estimator = DPTForDepthEstimation.from_pretrained(
172
+ "Intel/dpt-hybrid-midas"
173
+ ).to(device)
174
+ feature_extractor = DPTImageProcessor.from_pretrained("Intel/dpt-hybrid-midas")
175
+
176
+ image = feature_extractor(images=image, return_tensors="pt").pixel_values.to(device)
177
+ # MPS doesn't support autocast, so use 'cpu' for autocast when on MPS
178
+ from dw import get_autocast_device_type
179
+
180
+ autocast_device = get_autocast_device_type()
181
+ if autocast_device == "cuda":
182
+ with torch.no_grad(), torch.autocast(autocast_device):
183
+ depth_map = depth_estimator(image).predicted_depth
184
+ else:
185
+ # For MPS/CPU, don't use autocast
186
+ with torch.no_grad():
187
+ depth_map = depth_estimator(image).predicted_depth
188
+
189
+ depth_map = torch.nn.functional.interpolate(
190
+ depth_map.unsqueeze(1),
191
+ size=size,
192
+ mode="bicubic",
193
+ align_corners=False,
194
+ )
195
+ depth_min = torch.amin(depth_map, dim=[1, 2, 3], keepdim=True)
196
+ depth_max = torch.amax(depth_map, dim=[1, 2, 3], keepdim=True)
197
+ depth_map = (depth_map - depth_min) / (depth_max - depth_min)
198
+ image = torch.cat([depth_map] * 3, dim=1)
199
+
200
+ image = image.permute(0, 2, 3, 1).cpu().numpy()[0]
201
+ image = Image.fromarray((image * 255.0).clip(0, 255).astype(np.uint8))
202
+ return image
203
+
204
+
205
+ def image_to_segmentation(image):
206
+ """Semantic segmentation map from the UperNet ConvNeXt model, colored by class."""
207
+ from transformers import AutoImageProcessor, UperNetForSemanticSegmentation
208
+
209
+ image_processor = AutoImageProcessor.from_pretrained(
210
+ "openmmlab/upernet-convnext-small"
211
+ )
212
+ image_segmentor = UperNetForSemanticSegmentation.from_pretrained(
213
+ "openmmlab/upernet-convnext-small"
214
+ )
215
+ pixel_values = image_processor(image, return_tensors="pt").pixel_values
216
+ with torch.no_grad():
217
+ outputs = image_segmentor(pixel_values)
218
+ seg = image_processor.post_process_semantic_segmentation(
219
+ outputs, target_sizes=[image.size[::-1]]
220
+ )[0]
221
+ color_seg = np.zeros(
222
+ (seg.shape[0], seg.shape[1], 3), dtype=np.uint8
223
+ ) # height, width, 3
224
+ for label, color in enumerate(ada_palette):
225
+ color_seg[seg == label, :] = color
226
+ color_seg = color_seg.astype(np.uint8)
227
+ return Image.fromarray(color_seg)
228
+
229
+
230
+ def get_image_size(image):
231
+ """Return the image's width and height in pixels."""
232
+ return {"width": image.width, "height": image.height}
233
+
234
+
235
+ def crop_square(img: Image) -> Image:
236
+ """Crop the image to a centered square of its shorter side."""
237
+ # Determine the shortest side
238
+ min_side = min(img.width, img.height)
239
+
240
+ # Calculate the left and right crop positions for centering
241
+ left = (img.width - min_side) // 2
242
+ right = left + min_side
243
+
244
+ # Calculate the top and bottom crop positions for centering
245
+ top = (img.height - min_side) // 2
246
+ bottom = top + min_side
247
+
248
+ # Crop the image
249
+ img_cropped = img.crop((left, top, right, bottom))
250
+
251
+ return img_cropped
252
+
253
+
254
+ def resize_center_crop(img, height=768, width=768):
255
+ """Crop the image to its centered square and resize to width x height."""
256
+ output_size = (width, height)
257
+ W, H = img.size
258
+
259
+ # Calculate dimensions to crop to the center
260
+ new_dimension = min(W, H)
261
+ left = (W - new_dimension) / 2
262
+ top = (H - new_dimension) / 2
263
+ right = (W + new_dimension) / 2
264
+ bottom = (H + new_dimension) / 2
265
+
266
+ # Crop and resize
267
+ img = img.crop((left, top, right, bottom))
268
+ img = img.resize(output_size)
269
+
270
+ return img
271
+
272
+
273
+ def resize_rescale(image, height=768, width=768):
274
+ """Resize the image to width x height, ignoring its original aspect ratio."""
275
+ input_image = image.convert("RGB")
276
+ return input_image.resize((width, height))
277
+
278
+
279
+ def resize_resample(image, resolution=1024):
280
+ """Resize the image so its shorter side is `resolution`, rounded to a
281
+ multiple of 64, preserving aspect ratio."""
282
+ input_image = image.convert("RGB")
283
+ W, H = input_image.size
284
+ k = float(resolution) / min(H, W)
285
+ H *= k
286
+ W *= k
287
+ H = int(round(H / 64.0)) * 64
288
+ W = int(round(W / 64.0)) * 64
289
+
290
+ return input_image.resize((W, H), resample=Image.LANCZOS)
291
+
292
+
293
+ # Standard aspect ratios used by SDXL, Flux, and similar models.
294
+ # Each entry is (width_ratio, height_ratio).
295
+ _DEFAULT_RATIOS = [
296
+ (1, 1),
297
+ (4, 3),
298
+ (3, 4),
299
+ (3, 2),
300
+ (2, 3),
301
+ (16, 9),
302
+ (9, 16),
303
+ (21, 9),
304
+ (9, 21),
305
+ ]
306
+
307
+
308
+ def resize_bucket(image, resolution=1024, ratios=None, alignment=64):
309
+ """Resize image to the closest model-native aspect ratio bucket.
310
+
311
+ Picks the standard ratio closest to the input image's natural aspect
312
+ ratio, then scales to fit within the target resolution (based on the
313
+ short side) with dimensions aligned to `alignment` pixels.
314
+
315
+ Args:
316
+ image: PIL Image to resize.
317
+ resolution: Target size for the short side in pixels (default: 1024).
318
+ ratios: Optional list of [w, h] ratio pairs. Defaults to standard
319
+ ratios used by SDXL/Flux (1:1, 4:3, 3:2, 16:9, etc.).
320
+ alignment: Round dimensions to this multiple (default: 64).
321
+
322
+ Returns:
323
+ PIL Image resized to the bucketed dimensions.
324
+ """
325
+ input_image = image.convert("RGB")
326
+ W, H = input_image.size
327
+ input_ratio = W / H
328
+
329
+ bucket_ratios = ratios if ratios is not None else _DEFAULT_RATIOS
330
+
331
+ # Find the closest aspect ratio
332
+ best_ratio = min(
333
+ bucket_ratios,
334
+ key=lambda r: abs((r[0] / r[1]) - input_ratio),
335
+ )
336
+
337
+ wr, hr = best_ratio
338
+ bucket_ratio = wr / hr
339
+
340
+ # Scale so the short side matches resolution, then align
341
+ if bucket_ratio >= 1.0:
342
+ # Landscape or square: height is the short side
343
+ out_h = int(round(resolution / alignment)) * alignment
344
+ out_w = int(round((out_h * bucket_ratio) / alignment)) * alignment
345
+ else:
346
+ # Portrait: width is the short side
347
+ out_w = int(round(resolution / alignment)) * alignment
348
+ out_h = int(round((out_w / bucket_ratio) / alignment)) * alignment
349
+
350
+ return input_image.resize((out_w, out_h), resample=Image.LANCZOS)
351
+
352
+
353
+ def recenter_crop(
354
+ image, center_x=0.5, center_y=0.5, crop=1.0, width=None, height=None, fill="edge"
355
+ ):
356
+ """Re-frame an image around a chosen point, at a chosen scale.
357
+
358
+ Takes a square window `crop` of the shorter side across, centred on
359
+ (center_x, center_y) in normalised 0-1 coordinates, and resizes it to
360
+ width x height. Giving a series of images the same crop size and the same
361
+ centre - each one measured on its own subject - registers them: whatever
362
+ each picture is of, the chosen feature lands on the same pixel at the same
363
+ size. That is what lets a hard cut between two unrelated images read as one
364
+ continuous subject rather than as two pictures.
365
+
366
+ The window is allowed to run off the edge of the source, since a feature
367
+ near a border is exactly the case that needs moving furthest. `fill` says
368
+ what lies outside: "edge" replicates the border pixels, "reflect" and
369
+ "symmetric" mirror them back inward, and anything else is read as a PIL
370
+ colour name or tuple. Replication leaves visible streaks against a texture
371
+ and mirroring does not, so a subject sitting on sand, water or sky wants
372
+ "symmetric"; a subject on flat black wants the colour.
373
+ """
374
+ if crop <= 0:
375
+ raise ValueError(f"crop must be greater than zero, got {crop}")
376
+
377
+ image = image.convert("RGB")
378
+ source_width, source_height = image.size
379
+ side = int(round(crop * min(source_width, source_height)))
380
+
381
+ left = int(round(center_x * source_width - side / 2))
382
+ top = int(round(center_y * source_height - side / 2))
383
+
384
+ pad_left = max(0, -left)
385
+ pad_top = max(0, -top)
386
+ pad_right = max(0, left + side - source_width)
387
+ pad_bottom = max(0, top + side - source_height)
388
+
389
+ if pad_left or pad_top or pad_right or pad_bottom:
390
+ if fill in ("edge", "reflect", "symmetric"):
391
+ padded = Image.fromarray(
392
+ np.pad(
393
+ np.asarray(image),
394
+ ((pad_top, pad_bottom), (pad_left, pad_right), (0, 0)),
395
+ mode=fill,
396
+ )
397
+ )
398
+ else:
399
+ padded = Image.new(
400
+ "RGB",
401
+ (
402
+ source_width + pad_left + pad_right,
403
+ source_height + pad_top + pad_bottom,
404
+ ),
405
+ fill,
406
+ )
407
+ padded.paste(image, (pad_left, pad_top))
408
+ image = padded
409
+ left += pad_left
410
+ top += pad_top
411
+
412
+ window = image.crop((left, top, left + side, top + side))
413
+
414
+ return window.resize((width or side, height or side), Image.LANCZOS)
415
+
416
+
417
+ def strip_exif(image):
418
+ """Remove all EXIF and metadata from an image.
419
+
420
+ Creates a clean copy with pixel data only — no GPS coordinates,
421
+ camera info, timestamps, or other embedded metadata.
422
+
423
+ Args:
424
+ image: PIL Image to strip.
425
+
426
+ Returns:
427
+ PIL Image with all metadata removed.
428
+ """
429
+ clean = Image.new(image.mode, image.size)
430
+ clean.paste(image)
431
+ return clean
432
+
433
+
434
+ def add_watermark(
435
+ image,
436
+ text="AI Generated",
437
+ position="bottom-right",
438
+ opacity=128,
439
+ font_size=0,
440
+ margin=10,
441
+ color=None,
442
+ ):
443
+ """Add a visible text watermark to an image.
444
+
445
+ Args:
446
+ image: PIL Image to watermark.
447
+ text: Watermark text (default: "AI Generated").
448
+ position: Placement — "bottom-right", "bottom-left", "top-right",
449
+ "top-left", or "center" (default: "bottom-right").
450
+ opacity: Text opacity 0-255 (default: 128).
451
+ font_size: Font size in pixels. 0 = auto-scale to ~3% of image height.
452
+ margin: Pixel margin from edges (default: 10).
453
+ color: RGB tuple for text color (default: white).
454
+
455
+ Returns:
456
+ PIL Image with watermark applied.
457
+ """
458
+ from PIL import ImageDraw, ImageFont
459
+
460
+ base = image.convert("RGBA")
461
+ overlay = Image.new("RGBA", base.size, (0, 0, 0, 0))
462
+ draw = ImageDraw.Draw(overlay)
463
+
464
+ if color is None:
465
+ color = (255, 255, 255)
466
+ fill = (*color, int(opacity))
467
+
468
+ if font_size <= 0:
469
+ font_size = max(12, base.height // 30)
470
+
471
+ try:
472
+ font = ImageFont.truetype("Arial", font_size)
473
+ except (IOError, OSError):
474
+ font = ImageFont.load_default(size=font_size)
475
+
476
+ bbox = draw.textbbox((0, 0), text, font=font)
477
+ text_w = bbox[2] - bbox[0]
478
+ text_h = bbox[3] - bbox[1]
479
+
480
+ positions = {
481
+ "bottom-right": (base.width - text_w - margin, base.height - text_h - margin),
482
+ "bottom-left": (margin, base.height - text_h - margin),
483
+ "top-right": (base.width - text_w - margin, margin),
484
+ "top-left": (margin, margin),
485
+ "center": ((base.width - text_w) // 2, (base.height - text_h) // 2),
486
+ }
487
+ xy = positions.get(position, positions["bottom-right"])
488
+
489
+ draw.text(xy, text, font=font, fill=fill)
490
+
491
+ result = Image.alpha_composite(base, overlay)
492
+ return result.convert("RGB")
493
+
494
+
495
+ # ---------------------------------------------------------------------------
496
+ # process_image dispatch table
497
+ #
498
+ # Every handler has the uniform signature (image, device, kwargs) -> result,
499
+ # so process_image is just a lookup + call. Built once at import time from
500
+ # the detector spec tables above plus direct entries for the plain PIL/task
501
+ # functions.
502
+ # ---------------------------------------------------------------------------
503
+
504
+ _PROCESSORS = {
505
+ "get_image_size": lambda image, device, kwargs: get_image_size(image),
506
+ "add_border_and_mask": lambda image, device, kwargs: add_border_and_mask(
507
+ image, **kwargs
508
+ ),
509
+ "add_border_and_mask_with_size": lambda image, device, kwargs: (
510
+ add_border_and_mask_with_size(image, **kwargs)
511
+ ),
512
+ "remove_background": _remove_background_handler,
513
+ # Raw cv2 Canny at native resolution - see image_to_canny() docstring
514
+ # comment for how this differs from "canny" below.
515
+ "canny_cv": lambda image, device, kwargs: image_to_canny(image, **kwargs),
516
+ "segmentation": lambda image, device, kwargs: image_to_segmentation(image),
517
+ "zoe_depth": lambda image, device, kwargs: get_zoe_depth_map(image, device),
518
+ "depth": lambda image, device, kwargs: image_to_depth(image, device, **kwargs),
519
+ "depth_estimator_tensor": _depth_estimator_tensor_handler,
520
+ "depth_estimator": _depth_estimator_handler,
521
+ "resize_center_crop": lambda image, device, kwargs: resize_center_crop(
522
+ image, **kwargs
523
+ ),
524
+ "resize_resample": lambda image, device, kwargs: resize_resample(image, **kwargs),
525
+ "crop_square": lambda image, device, kwargs: crop_square(image, **kwargs),
526
+ "recenter_crop": lambda image, device, kwargs: recenter_crop(image, **kwargs),
527
+ "resize_rescale": lambda image, device, kwargs: resize_rescale(image, **kwargs),
528
+ "resize_bucket": lambda image, device, kwargs: resize_bucket(image, **kwargs),
529
+ "strip_exif": lambda image, device, kwargs: strip_exif(image),
530
+ "add_watermark": lambda image, device, kwargs: add_watermark(image, **kwargs),
531
+ }
532
+
533
+ for _name, (
534
+ _attr,
535
+ _repo_args,
536
+ _repo_kwargs,
537
+ _call_kwargs,
538
+ ) in _PRETRAINED_DETECTOR_SPECS.items():
539
+ _PROCESSORS[_name] = _make_pretrained_handler(
540
+ _attr, _repo_args, _repo_kwargs, _call_kwargs
541
+ )
542
+
543
+ # sam is the one from_pretrained detector never moved to device - matches
544
+ # the pre-refactor behavior, which called it straight off from_pretrained().
545
+ _PROCESSORS["sam"] = _make_pretrained_handler(
546
+ "SamDetector",
547
+ ("ybelkada/segment-anything",),
548
+ {"subfolder": "checkpoints"},
549
+ {},
550
+ to_device=False,
551
+ )
552
+
553
+ for _name, _attr in _ZERO_ARG_DETECTOR_SPECS.items():
554
+ _PROCESSORS[_name] = _make_zero_arg_handler(_attr)
555
+
556
+ _PROCESSORS["dw_pose"] = _dw_pose_handler
557
+
558
+ del _name, _attr, _repo_args, _repo_kwargs, _call_kwargs
559
+
560
+
561
+ def available_processors():
562
+ """Return the sorted list of processor names process_image accepts.
563
+
564
+ Used by command registration to enumerate supported image processors
565
+ without duplicating this dispatch table.
566
+ """
567
+ return sorted(_PROCESSORS)
568
+
569
+
570
+ # Processors whose handler is a plain (image, device, kwargs) -> function(image, **kwargs)
571
+ # forward - i.e. every argument beyond `image` is the named function's own, so
572
+ # introspection can read them straight off its signature and docstring instead of
573
+ # reporting the generic (image, device) shape every processor otherwise shares
574
+ # (#350). Detector-backed processors (controlnet_aux, transformers, dw_pose, sam)
575
+ # are left out on purpose: their real argument shape is the detector's __call__,
576
+ # not a Python function get_task can point at.
577
+ _PROCESSOR_TARGETS = {
578
+ "get_image_size": get_image_size,
579
+ "add_border_and_mask": add_border_and_mask,
580
+ "add_border_and_mask_with_size": add_border_and_mask_with_size,
581
+ "canny_cv": image_to_canny,
582
+ "segmentation": image_to_segmentation,
583
+ "resize_center_crop": resize_center_crop,
584
+ "resize_resample": resize_resample,
585
+ "crop_square": crop_square,
586
+ "recenter_crop": recenter_crop,
587
+ "resize_rescale": resize_rescale,
588
+ "resize_bucket": resize_bucket,
589
+ "strip_exif": strip_exif,
590
+ "add_watermark": add_watermark,
591
+ }
592
+
593
+
594
+ def image_processor_target(processor):
595
+ """The plain function backing `processor`'s handler, or None when the
596
+ processor is detector-backed and has no such function to introspect."""
597
+ return _PROCESSOR_TARGETS.get(processor)
598
+
599
+
600
+ def process_image(image, processor, device, kwargs):
601
+ processor = processor.lower()
602
+
603
+ handler = _PROCESSORS.get(processor)
604
+ if handler is None:
605
+ raise Exception(f"Unknown image processor type: {processor}")
606
+
607
+ return handler(image, device, kwargs)
608
+
609
+
610
+ ada_palette = np.asarray(
611
+ [
612
+ [0, 0, 0],
613
+ [120, 120, 120],
614
+ [180, 120, 120],
615
+ [6, 230, 230],
616
+ [80, 50, 50],
617
+ [4, 200, 3],
618
+ [120, 120, 80],
619
+ [140, 140, 140],
620
+ [204, 5, 255],
621
+ [230, 230, 230],
622
+ [4, 250, 7],
623
+ [224, 5, 255],
624
+ [235, 255, 7],
625
+ [150, 5, 61],
626
+ [120, 120, 70],
627
+ [8, 255, 51],
628
+ [255, 6, 82],
629
+ [143, 255, 140],
630
+ [204, 255, 4],
631
+ [255, 51, 7],
632
+ [204, 70, 3],
633
+ [0, 102, 200],
634
+ [61, 230, 250],
635
+ [255, 6, 51],
636
+ [11, 102, 255],
637
+ [255, 7, 71],
638
+ [255, 9, 224],
639
+ [9, 7, 230],
640
+ [220, 220, 220],
641
+ [255, 9, 92],
642
+ [112, 9, 255],
643
+ [8, 255, 214],
644
+ [7, 255, 224],
645
+ [255, 184, 6],
646
+ [10, 255, 71],
647
+ [255, 41, 10],
648
+ [7, 255, 255],
649
+ [224, 255, 8],
650
+ [102, 8, 255],
651
+ [255, 61, 6],
652
+ [255, 194, 7],
653
+ [255, 122, 8],
654
+ [0, 255, 20],
655
+ [255, 8, 41],
656
+ [255, 5, 153],
657
+ [6, 51, 255],
658
+ [235, 12, 255],
659
+ [160, 150, 20],
660
+ [0, 163, 255],
661
+ [140, 140, 140],
662
+ [250, 10, 15],
663
+ [20, 255, 0],
664
+ [31, 255, 0],
665
+ [255, 31, 0],
666
+ [255, 224, 0],
667
+ [153, 255, 0],
668
+ [0, 0, 255],
669
+ [255, 71, 0],
670
+ [0, 235, 255],
671
+ [0, 173, 255],
672
+ [31, 0, 255],
673
+ [11, 200, 200],
674
+ [255, 82, 0],
675
+ [0, 255, 245],
676
+ [0, 61, 255],
677
+ [0, 255, 112],
678
+ [0, 255, 133],
679
+ [255, 0, 0],
680
+ [255, 163, 0],
681
+ [255, 102, 0],
682
+ [194, 255, 0],
683
+ [0, 143, 255],
684
+ [51, 255, 0],
685
+ [0, 82, 255],
686
+ [0, 255, 41],
687
+ [0, 255, 173],
688
+ [10, 0, 255],
689
+ [173, 255, 0],
690
+ [0, 255, 153],
691
+ [255, 92, 0],
692
+ [255, 0, 255],
693
+ [255, 0, 245],
694
+ [255, 0, 102],
695
+ [255, 173, 0],
696
+ [255, 0, 20],
697
+ [255, 184, 184],
698
+ [0, 31, 255],
699
+ [0, 255, 61],
700
+ [0, 71, 255],
701
+ [255, 0, 204],
702
+ [0, 255, 194],
703
+ [0, 255, 82],
704
+ [0, 10, 255],
705
+ [0, 112, 255],
706
+ [51, 0, 255],
707
+ [0, 194, 255],
708
+ [0, 122, 255],
709
+ [0, 255, 163],
710
+ [255, 153, 0],
711
+ [0, 255, 10],
712
+ [255, 112, 0],
713
+ [143, 255, 0],
714
+ [82, 0, 255],
715
+ [163, 255, 0],
716
+ [255, 235, 0],
717
+ [8, 184, 170],
718
+ [133, 0, 255],
719
+ [0, 255, 92],
720
+ [184, 0, 255],
721
+ [255, 0, 31],
722
+ [0, 184, 255],
723
+ [0, 214, 255],
724
+ [255, 0, 112],
725
+ [92, 255, 0],
726
+ [0, 224, 255],
727
+ [112, 224, 255],
728
+ [70, 184, 160],
729
+ [163, 0, 255],
730
+ [153, 0, 255],
731
+ [71, 255, 0],
732
+ [255, 0, 163],
733
+ [255, 204, 0],
734
+ [255, 0, 143],
735
+ [0, 255, 235],
736
+ [133, 255, 0],
737
+ [255, 0, 235],
738
+ [245, 0, 255],
739
+ [255, 0, 122],
740
+ [255, 245, 0],
741
+ [10, 190, 212],
742
+ [214, 255, 0],
743
+ [0, 204, 255],
744
+ [20, 0, 255],
745
+ [255, 255, 0],
746
+ [0, 153, 255],
747
+ [0, 41, 255],
748
+ [0, 255, 204],
749
+ [41, 0, 255],
750
+ [41, 255, 0],
751
+ [173, 0, 255],
752
+ [0, 245, 255],
753
+ [71, 0, 255],
754
+ [122, 0, 255],
755
+ [0, 255, 184],
756
+ [0, 92, 255],
757
+ [184, 255, 0],
758
+ [0, 133, 255],
759
+ [255, 214, 0],
760
+ [25, 194, 194],
761
+ [102, 255, 0],
762
+ [92, 0, 255],
763
+ ]
764
+ )