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.
- diffusers_workflow-0.4.0.dist-info/METADATA +318 -0
- diffusers_workflow-0.4.0.dist-info/RECORD +260 -0
- diffusers_workflow-0.4.0.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0.dist-info/entry_points.txt +7 -0
- diffusers_workflow-0.4.0.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0.dist-info/top_level.txt +2 -0
- dw/__init__.py +440 -0
- dw/adapter_compatibility.py +226 -0
- dw/arguments.py +1231 -0
- dw/assessment_rules.py +159 -0
- dw/assets.py +130 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +146 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/content_types.py +150 -0
- dw/dissolve_frame_errors.py +121 -0
- dw/docs/ACCELERATION.md +352 -0
- dw/docs/AGENT_LOOP.md +95 -0
- dw/docs/DEPENDENCIES.md +91 -0
- dw/docs/IP_ADAPTER.md +109 -0
- dw/docs/LORAS.md +131 -0
- dw/docs/MCP.md +517 -0
- dw/docs/PROMPT_WEIGHTING.md +78 -0
- dw/docs/QUANTIZATION.md +230 -0
- dw/docs/RECIPES_24GB.md +201 -0
- dw/docs/RELEASING.md +195 -0
- dw/docs/REMOTE.md +140 -0
- dw/docs/REPL_COMMANDS.md +121 -0
- dw/docs/REPL_WORKER_GUIDE.md +51 -0
- dw/docs/SECURITY.md +272 -0
- dw/docs/SECURITY_QUICKREF.md +112 -0
- dw/docs/SERVER.md +679 -0
- dw/docs/TASKS.md +1741 -0
- dw/docs/TESTING.md +71 -0
- dw/docs/WORKFLOW_GUIDE.md +2038 -0
- dw/docs/WORKSPACES.md +316 -0
- dw/download_watch.py +335 -0
- dw/elision.py +306 -0
- dw/events.py +275 -0
- dw/for_each.py +409 -0
- dw/host_memory.py +258 -0
- dw/host_memory_projection.py +230 -0
- dw/hub_cache.py +432 -0
- dw/introspection.py +1228 -0
- dw/kernel_availability.py +208 -0
- dw/locations.py +599 -0
- dw/log_setup.py +45 -0
- dw/loudness.py +82 -0
- dw/media_audio.py +217 -0
- dw/media_frames.py +367 -0
- dw/media_info.py +297 -0
- dw/pipeline_processors/chain.py +821 -0
- dw/pipeline_processors/config_objects.py +237 -0
- dw/pipeline_processors/pipeline.py +2297 -0
- dw/pipeline_processors/remote.py +46 -0
- dw/plan.py +920 -0
- dw/previous_results.py +411 -0
- dw/probe_paths.py +59 -0
- dw/prompt_schema.json +48 -0
- dw/prompt_weighting.py +378 -0
- dw/prompts.py +159 -0
- dw/realize.py +250 -0
- dw/reference_limits.py +215 -0
- dw/reference_names.py +125 -0
- dw/repl.py +338 -0
- dw/repl_commands.py +836 -0
- dw/repl_worker.py +159 -0
- dw/result.py +1720 -0
- dw/result_fps.py +82 -0
- dw/run.py +162 -0
- dw/runs.py +768 -0
- dw/scalar_result_validation.py +97 -0
- dw/schema.py +283 -0
- dw/security.py +1038 -0
- dw/select_validation.py +115 -0
- dw/serve.py +277 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +4586 -0
- dw/server/assess.py +132 -0
- dw/server/catalog_shape.py +487 -0
- dw/server/enhancers.py +129 -0
- dw/server/exports.py +480 -0
- dw/server/guides.py +257 -0
- dw/server/jobs.py +1561 -0
- dw/server/mcp_mount.py +95 -0
- dw/server/netinfo.py +124 -0
- dw/server/observed_cost.py +379 -0
- dw/server/sysinfo.py +71 -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-CPznxfY8.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-CpWcotrd.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-CXtRM8N4.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-N7x-6NMY.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-PhsdjHSr.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-Dgj0SEok.js +1 -0
- dw/server/ui/assets/index-3Vw6WAPW.css +1 -0
- dw/server/ui/assets/index-DgrYhQd9.js +43 -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-BJqN9Qhv.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-DbM4SWSv.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-BWr8lEc4.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-C1iGuDyZ.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-DAdMi_0p.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-BGCeEqaw.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-Bcn70HdC.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-D1HmNnby.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-D6u0XmOW.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-BU6v-LMV.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-Cn7cTUKr.js +1 -0
- dw/server/ui/assets/xml--0LP2Lwk.js +1 -0
- dw/server/ui/assets/yaml-mpBg9jnt.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/server/updater.py +192 -0
- dw/settings.py +98 -0
- dw/shot_span_preflight.py +116 -0
- dw/shots.py +359 -0
- dw/slice_preflight.py +148 -0
- dw/step.py +187 -0
- dw/step_cache.py +442 -0
- dw/subfolders.py +107 -0
- dw/task_domains.py +307 -0
- dw/tasks/assess.py +826 -0
- dw/tasks/audio_transcription.py +88 -0
- dw/tasks/audio_utils.py +1862 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/compose_text.py +74 -0
- dw/tasks/concat_videos.py +300 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/dissolve_videos.py +342 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +173 -0
- dw/tasks/grade.py +97 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +764 -0
- dw/tasks/interpolate_frames.py +252 -0
- dw/tasks/judge.py +68 -0
- dw/tasks/model_cache.py +55 -0
- dw/tasks/pair_audio.py +268 -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/select.py +111 -0
- dw/tasks/speech_generation.py +228 -0
- dw/tasks/stabilize.py +129 -0
- dw/tasks/task.py +920 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +169 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +624 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +381 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +231 -0
- dw/validate.py +68 -0
- dw/variable_constraints.py +444 -0
- dw/variables.py +443 -0
- dw/video_extensions.py +141 -0
- dw/vram_estimate.py +116 -0
- dw/worker.py +764 -0
- dw/workflow.py +2007 -0
- dw/workflow_schema.json +1346 -0
- dw/workflow_sources.py +383 -0
- dw/workflows/h3_context_ir.json +57 -0
- dw/workflows/test.json +31 -0
- dw/workspace.py +730 -0
- dw_mcp/__init__.py +6 -0
- dw_mcp/__main__.py +133 -0
- dw_mcp/assets.py +336 -0
- dw_mcp/authoring.py +114 -0
- dw_mcp/catalog.py +360 -0
- dw_mcp/client.py +486 -0
- dw_mcp/diagnose.py +371 -0
- dw_mcp/exports.py +84 -0
- dw_mcp/guides.py +35 -0
- dw_mcp/media.py +638 -0
- dw_mcp/models.py +97 -0
- dw_mcp/prompts.py +104 -0
- dw_mcp/server.py +1343 -0
- dw_mcp/workspaces.py +212 -0
dw/tasks/tensor_image.py
ADDED
|
@@ -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
|