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,57 @@
1
+ """
2
+ Shared PIL <-> float tensor conversions.
3
+
4
+ Consolidates the PIL-to-tensor round trip that was previously hand-rolled
5
+ independently in upscale.py, restore_faces.py, and interpolate_frames.py.
6
+ """
7
+
8
+ import numpy as np
9
+ import torch
10
+ from PIL import Image
11
+
12
+
13
+ def pil_to_float_tensor(image, device, dtype=None):
14
+ """Convert a PIL image to a (1, 3, H, W) float tensor in [0, 1] on device.
15
+
16
+ The image is coerced to RGB first, so single-channel or RGBA inputs are
17
+ handled consistently. `dtype` defaults to float32 (the numpy source
18
+ precision); pass e.g. `torch.float16` to cast directly to a model's
19
+ working precision.
20
+
21
+ Args:
22
+ image: PIL Image
23
+ device: Target device (str or torch.device)
24
+ dtype: Optional torch dtype to cast to (default: float32)
25
+
26
+ Returns:
27
+ torch.Tensor of shape (1, 3, H, W)
28
+ """
29
+ arr = np.array(image.convert("RGB")).astype(np.float32) / 255.0
30
+ tensor = torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0).to(device)
31
+ if dtype is not None:
32
+ tensor = tensor.to(dtype=dtype)
33
+ return tensor
34
+
35
+
36
+ def float_tensor_to_pil(tensor):
37
+ """Convert a (1, 3, H, W) or (3, H, W) float tensor in [0, 1] to a PIL RGB image.
38
+
39
+ Quantizes to uint8 via rounding (`.round()`) rather than truncation
40
+ (a bare truncating cast), matching diffusers' VaeImageProcessor.numpy_to_pil
41
+ behavior. This matters: with truncation, an exact 8-bit value that
42
+ round-trips through [0, 1] float (e.g. 128/255) can land a hair below
43
+ its integer (127.999...) and get chopped down a level instead of
44
+ landing back on 128. Rounding fixes that, at the cost of at most 1/255
45
+ of drift per channel versus the old truncating behavior for values that
46
+ were never exact to begin with.
47
+
48
+ Args:
49
+ tensor: torch.Tensor of shape (1, 3, H, W) or (3, H, W), values in [0, 1]
50
+
51
+ Returns:
52
+ PIL.Image.Image in RGB mode
53
+ """
54
+ if tensor.dim() == 4:
55
+ tensor = tensor.squeeze(0)
56
+ arr = tensor.permute(1, 2, 0).mul(255).round().clamp(0, 255).byte().cpu().numpy()
57
+ return Image.fromarray(arr)
@@ -0,0 +1,169 @@
1
+ """
2
+ Text generation via HuggingFace transformers.
3
+
4
+ Takes a prompt (and optional system prompt) and generates text using a
5
+ local language model. Useful for prompt expansion, rewriting, and
6
+ other text-to-text tasks.
7
+
8
+ Supplying an image switches to a vision-language model, so the generated
9
+ text can describe what is actually in the picture rather than what the
10
+ prompt guesses is there. That is the difference between an image-conditioned
11
+ workflow whose prompt agrees with its keyframe and one whose prompt fights it.
12
+ """
13
+
14
+ import logging
15
+ from transformers import pipeline as hf_pipeline
16
+ from .. import preferred_task_dtype
17
+ from .model_cache import cached_model, hf_pipeline_placement
18
+
19
+ logger = logging.getLogger("dw")
20
+
21
+ _DEFAULT_MODEL = "Qwen/Qwen2.5-1.5B-Instruct"
22
+ # Small enough to stand in for the old captioning default; override for
23
+ # anything needing real detail
24
+ _DEFAULT_VISION_MODEL = "HuggingFaceTB/SmolVLM-256M-Instruct"
25
+
26
+ # Greedy decoding against a long, rigid format specification makes the vision
27
+ # models loop - finishing the answer, then repeating its closing sections until
28
+ # they run out of tokens. A penalty stops that without giving up reproducible
29
+ # output, which sampling would. Measured on Qwen3-VL against the H3 prompt
30
+ # spec: 1.05 still looped and filled the whole budget, 1.15 ended on its own at
31
+ # a length matching the format's own guidance. The text models do not need it
32
+ _VISION_REPETITION_PENALTY = 1.15
33
+
34
+
35
+ def _image_part(image):
36
+ """The chat content entry for an image.
37
+
38
+ A PIL image goes in as an object; a string is a URL or path the pipeline
39
+ loads itself, and it has to be declared as such - passing it under "image"
40
+ would hand the processor a bare string where it expects pixels.
41
+ """
42
+ return {"type": "image", "url" if isinstance(image, str) else "image": image}
43
+
44
+
45
+ def _build_messages(prompt, system_prompt, image):
46
+ """Chat messages in the shape the chosen pipeline expects.
47
+
48
+ Text-generation models take plain string content. Vision models take a
49
+ list of typed parts, because the image is a part of the message rather
50
+ than something alongside it.
51
+ """
52
+ messages = []
53
+ if image is None:
54
+ if system_prompt is not None:
55
+ messages.append({"role": "system", "content": system_prompt})
56
+ messages.append({"role": "user", "content": prompt})
57
+ else:
58
+ if system_prompt is not None:
59
+ messages.append(
60
+ {"role": "system", "content": [{"type": "text", "text": system_prompt}]}
61
+ )
62
+ messages.append(
63
+ {
64
+ "role": "user",
65
+ "content": [_image_part(image), {"type": "text", "text": prompt}],
66
+ }
67
+ )
68
+ return messages
69
+
70
+
71
+ def generate_text(prompt, device="cpu", **kwargs):
72
+ """Generate text from a prompt using a local language model.
73
+
74
+ Args:
75
+ prompt: The user message / prompt to expand or transform.
76
+ device: Target device ("cuda", "mps", "cpu").
77
+ **kwargs:
78
+ model_name: HuggingFace model ID. Defaults to
79
+ Qwen/Qwen2.5-1.5B-Instruct, or a small vision-language model
80
+ when an image is supplied. An image needs a model that can
81
+ accept one - a text-only model will fail to load as one.
82
+ system_prompt: Optional system instruction for the model.
83
+ max_new_tokens: Max tokens to generate (default: 500).
84
+ image: Optional PIL image, URL or path. Its presence is what
85
+ selects the vision pipeline.
86
+ repetition_penalty: Vision pipeline only (default: 1.15). Raise it
87
+ if a model still repeats itself, or set 1.0 to disable.
88
+ generate_kwargs: Anything else to hand the model's generate() -
89
+ no_repeat_ngram_size, top_p, min_new_tokens and so on. Merged
90
+ over what this function sets, so it can override those too.
91
+
92
+ Returns:
93
+ Generated text string.
94
+ """
95
+ # A workflow declaring an optional image passes it through as null when the
96
+ # caller supplies none, and an empty string is the same statement
97
+ image = kwargs.get("image", None) or None
98
+ system_prompt = kwargs.get("system_prompt", None)
99
+ max_new_tokens = int(kwargs.get("max_new_tokens", 500))
100
+
101
+ if image is None:
102
+ pipeline_task = "text-generation"
103
+ model_name = kwargs.get("model_name", _DEFAULT_MODEL)
104
+ else:
105
+ pipeline_task = "image-text-to-text"
106
+ model_name = kwargs.get("model_name", _DEFAULT_VISION_MODEL)
107
+
108
+ dtype = preferred_task_dtype(device)
109
+
110
+ def load_pipe():
111
+ logger.info(f"Generating text with {model_name} on {device}")
112
+ placement = hf_pipeline_placement(device)
113
+ return hf_pipeline(
114
+ pipeline_task,
115
+ model=model_name,
116
+ torch_dtype=dtype,
117
+ **placement,
118
+ )
119
+
120
+ # The task is part of the identity - the same model name can be loaded
121
+ # under either pipeline, and they are not interchangeable
122
+ pipe = cached_model(
123
+ ("text_generation", pipeline_task, model_name, str(device), str(dtype)),
124
+ load_pipe,
125
+ )
126
+
127
+ messages = _build_messages(prompt, system_prompt, image)
128
+
129
+ # Decoding is greedy, so the same input returns the same text every run -
130
+ # which is what a workflow wants, and why there is nothing here to seed
131
+ generation = {"do_sample": False}
132
+ if image is not None:
133
+ generation["repetition_penalty"] = float(
134
+ kwargs.get("repetition_penalty", _VISION_REPETITION_PENALTY)
135
+ )
136
+ # Last, so a workflow can override anything decided above
137
+ generation.update(kwargs.get("generate_kwargs") or {})
138
+
139
+ return _generate(pipe, messages, image, max_new_tokens, generation)
140
+
141
+
142
+ def _generate(pipe, messages, image, max_new_tokens, generation):
143
+ """Run the pipeline, which each path calls differently."""
144
+ if image is None:
145
+ # This pipeline collects anything it does not name into the arguments it
146
+ # forwards to generate(), so settings go in as plain keywords
147
+ results = pipe(
148
+ messages,
149
+ max_new_tokens=max_new_tokens,
150
+ return_full_text=False,
151
+ **generation,
152
+ )
153
+ else:
154
+ # Images live inside the messages, so the chat goes in as `text` - the
155
+ # pipeline rejects a chat and an `images` argument together. Generation
156
+ # settings have to go through generate_kwargs here: anything else this
157
+ # pipeline does not name explicitly is forwarded to the processor and
158
+ # dropped, so a bare do_sample=False would leave sampling on. Passing
159
+ # max_new_tokens both ways is an error, so it stays a direct argument
160
+ results = pipe(
161
+ text=messages,
162
+ max_new_tokens=max_new_tokens,
163
+ return_full_text=False,
164
+ generate_kwargs=generation,
165
+ )
166
+
167
+ text = results[0]["generated_text"].strip()
168
+ logger.info(f"Generated: {text[:100]}{'...' if len(text) > 100 else ''}")
169
+ return text
@@ -0,0 +1,80 @@
1
+ """
2
+ Reduce a generated block of text to a known set of labelled sections.
3
+
4
+ A language model asked for a rigid format usually produces it and then keeps
5
+ going - restating the description, appending a summary, or looping until it
6
+ runs out of tokens. Downstream that trailing text is not free: a prompt is
7
+ conditioning, and a pipeline that does not truncate spends memory and attention
8
+ on whatever arrived. Rather than trying to talk the model out of it, keep the
9
+ parts that were asked for and drop the rest.
10
+ """
11
+
12
+ import logging
13
+ import re
14
+
15
+ logger = logging.getLogger("dw")
16
+
17
+
18
+ def extract_sections(text, sections, keep_preamble=True):
19
+ """Keep only the named sections, once each, in the order they are declared.
20
+
21
+ A section runs from its `label:` to the end of that paragraph - the format
22
+ these prompts use puts each field in one continuous paragraph, so a blank
23
+ line ends it. Anything outside a named section is dropped, which is what
24
+ removes a trailing restatement whether or not it repeats verbatim.
25
+
26
+ Args:
27
+ text: The generated text.
28
+ sections: Section labels to keep, in the order they should appear.
29
+ keep_preamble: Keep any text before the first label. These formats put
30
+ an instruction line above the fields, which is part of the output.
31
+
32
+ Returns:
33
+ The reassembled text.
34
+ """
35
+ if not sections:
36
+ return text.strip()
37
+
38
+ # Matched case-insensitively because the models capitalise labels however
39
+ # they please - one writes overall_soundscape, another Overall_soundscape,
40
+ # and dropping a section over its first letter would lose real content. The
41
+ # declared spelling is what gets written back out
42
+ labels = "|".join(re.escape(name) for name in sections)
43
+ matches = list(re.finditer(rf"^({labels}):", text, re.M | re.I))
44
+ canonical = {name.lower(): name for name in sections}
45
+
46
+ if not matches:
47
+ logger.warning(
48
+ "No sections of %s found - leaving the text as it is", list(sections)
49
+ )
50
+ return text.strip()
51
+
52
+ bodies = {}
53
+ for index, match in enumerate(matches):
54
+ name = canonical[match.group(1).lower()]
55
+ if name in bodies:
56
+ # A repeat is the model starting over, not new content
57
+ continue
58
+ # The body ends at the next label or the end of its paragraph, whichever
59
+ # comes first. A label may sit on its own line, so leading blank space
60
+ # is skipped before looking for the break that ends it
61
+ end = matches[index + 1].start() if index + 1 < len(matches) else len(text)
62
+ body = text[match.end() : end]
63
+ stripped = body.lstrip()
64
+ offset = len(body) - len(stripped)
65
+ paragraph = re.search(r"\n\s*\n", body[offset:])
66
+ if paragraph:
67
+ body = body[offset : offset + paragraph.start()]
68
+ bodies[name] = body.strip()
69
+
70
+ kept = [f"{name}: {bodies[name]}" for name in sections if name in bodies]
71
+
72
+ preamble = text[: matches[0].start()].strip()
73
+ if keep_preamble and preamble:
74
+ kept.insert(0, preamble)
75
+
76
+ result = "\n\n".join(kept)
77
+ dropped = len(text.strip()) - len(result)
78
+ if dropped > 0:
79
+ logger.info(f"Trimmed {dropped} characters outside the requested sections")
80
+ return result
dw/tasks/upscale.py ADDED
@@ -0,0 +1,203 @@
1
+ """
2
+ Image upscaling via spandrel.
3
+
4
+ Supports 40+ super-resolution architectures (ESRGAN, SwinIR, HAT, DAT, etc.)
5
+ with automatic model detection from .pth/.safetensors files.
6
+
7
+ Models can be loaded from local files or HuggingFace Hub repos.
8
+ """
9
+
10
+ import logging
11
+ import torch
12
+
13
+ from .model_cache import cached_model
14
+ from .tensor_image import pil_to_float_tensor, float_tensor_to_pil
15
+
16
+ logger = logging.getLogger("dw")
17
+
18
+ # Maximum tile size before we switch to tiled processing
19
+ _MAX_PIXELS_NO_TILE = 512 * 512
20
+
21
+
22
+ def upscale_image(image, model_name, device="cpu", **kwargs):
23
+ """Upscale an image using a spandrel-compatible super-resolution model.
24
+
25
+ Args:
26
+ image: PIL Image to upscale
27
+ model_name: HuggingFace repo ID (e.g., "user/repo") with optional filename,
28
+ or local path to .pth/.safetensors file.
29
+ device: Target device ("cuda", "mps", "cpu")
30
+ **kwargs:
31
+ filename: Weight file name within a HF repo (default: auto-detect)
32
+ tile_size: Tile size for large images (default: 512)
33
+ tile_overlap: Overlap between tiles in pixels (default: 32)
34
+
35
+ Returns:
36
+ PIL Image (upscaled)
37
+ """
38
+ try:
39
+ from spandrel import ModelLoader, ImageModelDescriptor
40
+ except ImportError:
41
+ raise ImportError(
42
+ "spandrel is required for upscaling. Install with: pip install spandrel"
43
+ )
44
+
45
+ filename = kwargs.get("filename", None)
46
+ tile_size = kwargs.get("tile_size", 512)
47
+ tile_overlap = kwargs.get("tile_overlap", 32)
48
+
49
+ def load_descriptor():
50
+ model_path = _resolve_model_path(model_name, filename)
51
+
52
+ logger.info(f"Loading upscale model from {model_path}")
53
+ loader = ModelLoader(device=torch.device(device))
54
+ result = loader.load_from_file(model_path)
55
+
56
+ if not isinstance(result, ImageModelDescriptor):
57
+ raise ValueError(
58
+ f"Model is not an image model (got {type(result).__name__}). "
59
+ f"Only image super-resolution models are supported."
60
+ )
61
+
62
+ logger.info(
63
+ f"Loaded {result.architecture.name} "
64
+ f"(scale: {result.scale}x, "
65
+ f"input: {result.input_channels}ch, "
66
+ f"output: {result.output_channels}ch)"
67
+ )
68
+
69
+ # Use half precision if supported and on GPU
70
+ if device != "cpu" and result.supports_half:
71
+ result.model.half()
72
+
73
+ return result
74
+
75
+ descriptor = cached_model(
76
+ ("upscale", model_name, filename, str(device)), load_descriptor
77
+ )
78
+ model_dtype = (
79
+ torch.float16
80
+ if (device != "cpu" and descriptor.supports_half)
81
+ else torch.float32
82
+ )
83
+
84
+ # Convert PIL to tensor
85
+ tensor = pil_to_float_tensor(image, device, dtype=model_dtype)
86
+
87
+ h, w = tensor.shape[2], tensor.shape[3]
88
+
89
+ if h * w <= _MAX_PIXELS_NO_TILE:
90
+ logger.debug(f"Upscaling {w}x{h} directly")
91
+ with torch.inference_mode():
92
+ output = descriptor(tensor)
93
+ else:
94
+ logger.debug(f"Upscaling {w}x{h} with {tile_size}px tiles")
95
+ output = _tiled_inference(descriptor, tensor, tile_size, tile_overlap)
96
+
97
+ # Convert back to PIL
98
+ result = float_tensor_to_pil(output)
99
+
100
+ logger.info(f"Upscaled {w}x{h} -> {result.width}x{result.height}")
101
+ return result
102
+
103
+
104
+ def _resolve_model_path(model_name, filename=None):
105
+ """Resolve a model name to a local file path.
106
+
107
+ Supports:
108
+ - Local file paths: "/path/to/model.pth"
109
+ - HuggingFace Hub: "user/repo" (auto-detects .pth/.safetensors)
110
+ - HuggingFace Hub with filename: model_name="user/repo", filename="model_x4.pth"
111
+ """
112
+ import os
113
+
114
+ # Local file
115
+ if os.path.exists(model_name):
116
+ return model_name
117
+
118
+ # HuggingFace Hub
119
+ try:
120
+ from huggingface_hub import hf_hub_download
121
+ except ImportError:
122
+ raise ImportError(
123
+ f"Model '{model_name}' is not a local file. "
124
+ "Install huggingface_hub to download from HF Hub."
125
+ )
126
+
127
+ if filename is not None:
128
+ logger.debug(f"Downloading {filename} from {model_name}")
129
+ return hf_hub_download(repo_id=model_name, filename=filename)
130
+
131
+ # Auto-detect: list repo files and find a model file
132
+ from huggingface_hub import list_repo_files
133
+
134
+ model_extensions = {".pth", ".pt", ".ckpt", ".safetensors"}
135
+ try:
136
+ files = list_repo_files(model_name)
137
+ except Exception as e:
138
+ raise ValueError(f"Could not access HuggingFace repo '{model_name}': {e}")
139
+
140
+ model_files = [f for f in files if any(f.endswith(ext) for ext in model_extensions)]
141
+ if not model_files:
142
+ raise ValueError(
143
+ f"No model files ({', '.join(model_extensions)}) found in '{model_name}'. "
144
+ f"Specify 'filename' explicitly."
145
+ )
146
+ if len(model_files) > 1:
147
+ raise ValueError(
148
+ f"Multiple model files found in '{model_name}': {model_files}. "
149
+ f"Specify 'filename' explicitly."
150
+ )
151
+
152
+ logger.debug(f"Downloading {model_files[0]} from {model_name}")
153
+ return hf_hub_download(repo_id=model_name, filename=model_files[0])
154
+
155
+
156
+ def _tiled_inference(descriptor, tensor, tile_size, overlap):
157
+ """Run model inference on overlapping tiles and blend results.
158
+
159
+ Splits the input into tiles, runs each through the model, then
160
+ blends overlapping regions with linear interpolation.
161
+ """
162
+ scale = descriptor.scale
163
+ _, c, h, w = tensor.shape
164
+ out_h, out_w = h * scale, w * scale
165
+ output = torch.zeros(1, c, out_h, out_w, device=tensor.device, dtype=tensor.dtype)
166
+ weight = torch.zeros(1, 1, out_h, out_w, device=tensor.device, dtype=tensor.dtype)
167
+
168
+ # Generate tile positions
169
+ y_positions = list(range(0, h, tile_size - overlap))
170
+ x_positions = list(range(0, w, tile_size - overlap))
171
+
172
+ # Clamp last tile to image boundary
173
+ y_positions = [min(y, max(0, h - tile_size)) for y in y_positions]
174
+ x_positions = [min(x, max(0, w - tile_size)) for x in x_positions]
175
+
176
+ # Deduplicate
177
+ y_positions = sorted(set(y_positions))
178
+ x_positions = sorted(set(x_positions))
179
+
180
+ total_tiles = len(y_positions) * len(x_positions)
181
+ logger.debug(f"Processing {total_tiles} tiles ({tile_size}px, {overlap}px overlap)")
182
+
183
+ tile_num = 0
184
+ for y in y_positions:
185
+ for x in x_positions:
186
+ tile_num += 1
187
+ th = min(tile_size, h - y)
188
+ tw = min(tile_size, w - x)
189
+
190
+ tile = tensor[:, :, y : y + th, x : x + tw]
191
+
192
+ with torch.inference_mode():
193
+ tile_out = descriptor(tile)
194
+
195
+ oy, ox = y * scale, x * scale
196
+ oth, otw = th * scale, tw * scale
197
+
198
+ output[:, :, oy : oy + oth, ox : ox + otw] += tile_out
199
+ weight[:, :, oy : oy + oth, ox : ox + otw] += 1
200
+
201
+ # Average overlapping regions
202
+ output = output / weight.clamp(min=1)
203
+ return output