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,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
@@ -0,0 +1,154 @@
1
+ """Frame access for videos in any of the shapes results carry them.
2
+
3
+ A video artifact can be a list of PIL images, a numpy array of frames
4
+ (frames, height, width, channels), a torch tensor (frames first, channels
5
+ first or last), or an AudioVideo pairing frames with their generated
6
+ soundtrack. extract_frame gives tasks and the segment-chaining loop one way
7
+ to pull a single frame out of any of them, always as a PIL image.
8
+ """
9
+
10
+ import numpy
11
+ import torch
12
+ from PIL import Image
13
+
14
+ from ..result import AudioVideo
15
+
16
+
17
+ def process_video(video, processor, device, kwargs):
18
+ processor = processor.lower()
19
+
20
+ if processor == "get_frame":
21
+ return get_frame(video, kwargs.get("frame_index", 0))
22
+
23
+ if processor == "get_last_frame":
24
+ return get_frame(video, -1)
25
+
26
+ if processor == "get_first_frame":
27
+ return get_frame(video, 0)
28
+
29
+ raise Exception(f"Unknown video processor type: {processor}")
30
+
31
+
32
+ def get_frame(video, frame_index=0):
33
+ return extract_frame(video, frame_index)
34
+
35
+
36
+ def extract_frame(video, index):
37
+ """Pull one frame out of a video, whatever the video's in-memory shape.
38
+
39
+ Args:
40
+ video: List of PIL images, numpy array or torch tensor of frames,
41
+ an AudioVideo, or a one-video batch wrapping any of those
42
+ index: Frame to extract; negative indexes count from the end
43
+
44
+ Returns:
45
+ The frame as a PIL image. Frames that already are PIL images are
46
+ returned as-is, not copied.
47
+ """
48
+ return _to_pil(_frames_of(video)[index])
49
+
50
+
51
+ def frame_count(video):
52
+ """Number of frames in a video of any supported shape."""
53
+ return len(_frames_of(video))
54
+
55
+
56
+ def frames_as_pil_list(video):
57
+ """The video's frames as a list of PIL images.
58
+
59
+ Frames that already are PIL images are carried over by identity; array and
60
+ tensor frames are converted the way extract_frame converts them.
61
+ """
62
+ return [_to_pil(frame) for frame in _frames_of(video)]
63
+
64
+
65
+ def frames_as_array(video):
66
+ """The video's frames as one (frames, height, width, channels) uint8 array.
67
+
68
+ The shape an argument that takes frames rather than a video wants - LTX-2's
69
+ keyframe conditions and IC-LoRA references, which the workflow hands what an
70
+ earlier step generated. One array is also one artifact, where a list of frames
71
+ would become one artifact per frame and multiply the step that consumed it.
72
+
73
+ Frames that already are a channels-last RGB array are converted in a single
74
+ operation; anything else goes through the same per-frame conversion
75
+ extract_frame uses.
76
+ """
77
+ frames = _frames_of(video)
78
+
79
+ if isinstance(frames, numpy.ndarray) and frames.ndim == 4 and frames.shape[-1] == 3:
80
+ if frames.dtype == numpy.uint8:
81
+ return frames
82
+ # Float frames are [0, 1] - diffusers' np output convention
83
+ return (numpy.clip(frames, 0.0, 1.0) * 255).round().astype(numpy.uint8)
84
+
85
+ return numpy.stack(
86
+ [numpy.asarray(_to_pil(frame).convert("RGB")) for frame in frames]
87
+ )
88
+
89
+
90
+ def _frames_of(video):
91
+ """Unwrap containers until an indexable run of frames remains."""
92
+ if isinstance(video, AudioVideo):
93
+ return _frames_of(video.frames)
94
+
95
+ if isinstance(video, list):
96
+ # A one-video batch - [[frame, ...]] or [ndarray] - unwraps to the video;
97
+ # a single-frame video - [frame] - is already the frames
98
+ if len(video) == 1 and not _is_frame(video[0]):
99
+ return _frames_of(video[0])
100
+ return video
101
+
102
+ if isinstance(video, numpy.ndarray):
103
+ if video.ndim == 3: # a lone frame
104
+ return video[numpy.newaxis, ...]
105
+ if video.ndim == 5 and video.shape[0] == 1: # a one-video batch
106
+ return video[0]
107
+ return video
108
+
109
+ if torch.is_tensor(video):
110
+ tensor = video.detach().cpu()
111
+ if tensor.ndim == 5 and tensor.shape[0] == 1: # a one-video batch
112
+ tensor = tensor[0]
113
+ if tensor.ndim == 3: # a lone frame
114
+ tensor = tensor.unsqueeze(0)
115
+ return tensor
116
+
117
+ raise TypeError(f"Cannot extract frames from a {type(video).__name__}")
118
+
119
+
120
+ def _is_frame(item):
121
+ """A single image: PIL, or a 3-dim array/tensor (height, width, channels)."""
122
+ if isinstance(item, Image.Image):
123
+ return True
124
+ if isinstance(item, numpy.ndarray) or torch.is_tensor(item):
125
+ return item.ndim == 3
126
+ return False
127
+
128
+
129
+ def _to_pil(frame):
130
+ """Convert one frame to a PIL image; PIL frames pass through untouched."""
131
+ if isinstance(frame, Image.Image):
132
+ return frame
133
+
134
+ if torch.is_tensor(frame):
135
+ frame = frame.detach().cpu().float().numpy()
136
+
137
+ if isinstance(frame, numpy.ndarray):
138
+ if frame.ndim != 3:
139
+ raise ValueError(f"A frame must have 3 dimensions, got {frame.ndim}")
140
+
141
+ # Channels-first (C, H, W) -> channels-last, the layout PIL expects
142
+ if frame.shape[0] in (1, 3, 4) and frame.shape[-1] not in (1, 3, 4):
143
+ frame = numpy.moveaxis(frame, 0, -1)
144
+
145
+ if frame.dtype != numpy.uint8:
146
+ # Float frames are [0, 1] - diffusers' np output convention
147
+ frame = (numpy.clip(frame, 0.0, 1.0) * 255).round().astype(numpy.uint8)
148
+
149
+ if frame.shape[-1] == 1: # grayscale
150
+ frame = frame[..., 0]
151
+
152
+ return Image.fromarray(frame)
153
+
154
+ raise TypeError(f"Cannot convert a {type(frame).__name__} to an image")
dw/tasks/zoe_depth.py ADDED
@@ -0,0 +1,71 @@
1
+ import torch
2
+ import matplotlib
3
+ import matplotlib.cm
4
+ import numpy as np
5
+ from PIL import Image
6
+
7
+ from .model_cache import cached_model
8
+
9
+
10
+ def load_zoe(device="cuda"):
11
+ def load():
12
+ torch.hub.help(
13
+ "intel-isl/MiDaS", "DPT_BEiT_L_384"
14
+ ) # Triggers fresh download of MiDaS repo
15
+ model_zoe_n = torch.hub.load(
16
+ "isl-org/ZoeDepth", "ZoeD_NK", pretrained=True
17
+ ).eval()
18
+ return model_zoe_n.to(device)
19
+
20
+ return cached_model(("zoe_depth", str(device)), load)
21
+
22
+
23
+ def colorize(
24
+ value,
25
+ vmin=None,
26
+ vmax=None,
27
+ cmap="gray_r",
28
+ invalid_val=-99,
29
+ invalid_mask=None,
30
+ background_color=(128, 128, 128, 255),
31
+ gamma_corrected=False,
32
+ value_transform=None,
33
+ ):
34
+ if isinstance(value, torch.Tensor):
35
+ value = value.detach().cpu().numpy()
36
+
37
+ value = value.squeeze()
38
+ if invalid_mask is None:
39
+ invalid_mask = value == invalid_val
40
+ mask = np.logical_not(invalid_mask)
41
+
42
+ # normalize
43
+ vmin = np.percentile(value[mask], 2) if vmin is None else vmin
44
+ vmax = np.percentile(value[mask], 85) if vmax is None else vmax
45
+ if vmin != vmax:
46
+ value = (value - vmin) / (vmax - vmin) # vmin..vmax
47
+ else:
48
+ # Avoid 0-division
49
+ value = value * 0.0
50
+
51
+ # squeeze last dim if it exists
52
+ # grey out the invalid values
53
+
54
+ value[invalid_mask] = np.nan
55
+ cmapper = matplotlib.cm.get_cmap(cmap)
56
+ if value_transform:
57
+ value = value_transform(value)
58
+ # value = value / value.max()
59
+ value = cmapper(value, bytes=True) # (nxmx4)
60
+
61
+ # img = value[:, :, :]
62
+ img = value[...]
63
+ img[invalid_mask] = background_color
64
+
65
+ # gamma correction
66
+ img = img / 255
67
+ img = np.power(img, 2.2)
68
+ img = img * 255
69
+ img = img.astype(np.uint8)
70
+ img = Image.fromarray(img)
71
+ return img