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.
- diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
- diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
- diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
- diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
- dw/__init__.py +353 -0
- dw/arguments.py +906 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +145 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/events.py +78 -0
- dw/hub_cache.py +289 -0
- dw/introspection.py +458 -0
- dw/log_setup.py +45 -0
- dw/pipeline_processors/chain.py +750 -0
- dw/pipeline_processors/config_objects.py +235 -0
- dw/pipeline_processors/pipeline.py +1687 -0
- dw/pipeline_processors/remote.py +18 -0
- dw/previous_results.py +259 -0
- dw/prompt_weighting.py +378 -0
- dw/repl.py +298 -0
- dw/repl_commands.py +808 -0
- dw/repl_worker.py +129 -0
- dw/result.py +850 -0
- dw/run.py +92 -0
- dw/schema.py +24 -0
- dw/security.py +379 -0
- dw/serve.py +70 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +588 -0
- dw/server/jobs.py +547 -0
- dw/server/ui/assets/abap-08VXUWAP.js +1 -0
- dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
- dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
- dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
- dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
- dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
- dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
- dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
- dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
- dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
- dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
- dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
- dw/server/ui/assets/css-DIMkf-bt.js +3 -0
- dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
- dw/server/ui/assets/cssMode-CEh6hWi2.js +1 -0
- dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
- dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
- dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
- dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
- dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
- dw/server/ui/assets/editor.api-CExg3_mM.js +847 -0
- dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
- dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
- dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
- dw/server/ui/assets/freemarker2-DH6orYh2.js +3 -0
- dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
- dw/server/ui/assets/go-C-y9NEjX.js +1 -0
- dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
- dw/server/ui/assets/handlebars-CbrMVW4Q.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-YDNPZw2M.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
- dw/server/ui/assets/index-B7-VcYS-.css +1 -0
- dw/server/ui/assets/index-D_EiPU3b.js +13 -0
- dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
- dw/server/ui/assets/java-BEtHBSE6.js +1 -0
- dw/server/ui/assets/javascript-dYuBvioq.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-CUqLM39V.js +7 -0
- dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
- dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
- dw/server/ui/assets/less-B9JPFI3C.js +2 -0
- dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
- dw/server/ui/assets/liquid-D6vxBzMv.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.js +4 -0
- dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
- dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
- dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
- dw/server/ui/assets/mdx-SHQb6vmD.js +1 -0
- dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
- dw/server/ui/assets/monaco--ixms01u.css +1 -0
- dw/server/ui/assets/monaco-CP-s5rcP.js +56 -0
- dw/server/ui/assets/msdax-DauUninz.js +1 -0
- dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
- dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
- dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
- dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
- dw/server/ui/assets/perl-oz_6vUea.js +1 -0
- dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
- dw/server/ui/assets/php-nr791fC2.js +1 -0
- dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
- dw/server/ui/assets/postiats-43DmfD33.js +1 -0
- dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
- dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
- dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
- dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
- dw/server/ui/assets/python-x0_EGHq9.js +1 -0
- dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
- dw/server/ui/assets/r-BwWrilGY.js +1 -0
- dw/server/ui/assets/razor-BZC4LQDP.js +1 -0
- dw/server/ui/assets/redis-ClamHrr6.js +1 -0
- dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
- dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
- dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
- dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
- dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
- dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
- dw/server/ui/assets/scheme-BeGwcela.js +1 -0
- dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
- dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
- dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
- dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
- dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
- dw/server/ui/assets/sql-NEE52Syq.js +1 -0
- dw/server/ui/assets/st-DbInun42.js +1 -0
- dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
- dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
- dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
- dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
- dw/server/ui/assets/tsMode-BTfA6SbD.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-CWA4MsNk.js +1 -0
- dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
- dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
- dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
- dw/server/ui/assets/workers-CWU0uvj5.js +1 -0
- dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
- dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/settings.py +77 -0
- dw/step.py +132 -0
- dw/tasks/audio_utils.py +266 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/concat_videos.py +80 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +139 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +661 -0
- dw/tasks/interpolate_frames.py +227 -0
- dw/tasks/model_cache.py +39 -0
- dw/tasks/pair_audio.py +58 -0
- dw/tasks/qr_code.py +19 -0
- dw/tasks/restore_faces.py +175 -0
- dw/tasks/rife_model.py +192 -0
- dw/tasks/segment.py +121 -0
- dw/tasks/task.py +474 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +168 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +154 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +376 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +68 -0
- dw/validate.py +43 -0
- dw/variables.py +153 -0
- dw/worker.py +517 -0
- dw/workflow.py +553 -0
- dw/workflow_schema.json +1157 -0
- dw/workflows/augment_prompt.json +65 -0
- dw/workflows/describe_image.json +58 -0
- dw/workflows/h3_context_ir.json +57 -0
- 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
|
dw/tasks/video_utils.py
ADDED
|
@@ -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
|