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
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
from PIL import Image
|
|
2
|
+
import torch
|
|
3
|
+
from torchvision import transforms
|
|
4
|
+
from transformers import AutoModelForImageSegmentation
|
|
5
|
+
|
|
6
|
+
from .model_cache import cached_model
|
|
7
|
+
|
|
8
|
+
_MODEL_NAME = "briaai/RMBG-2.0"
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def remove_background(image: Image, device) -> Image:
|
|
12
|
+
# Model settings
|
|
13
|
+
def load_model():
|
|
14
|
+
model = AutoModelForImageSegmentation.from_pretrained(
|
|
15
|
+
_MODEL_NAME, trust_remote_code=True
|
|
16
|
+
)
|
|
17
|
+
model.to(device)
|
|
18
|
+
model.eval()
|
|
19
|
+
return model
|
|
20
|
+
|
|
21
|
+
model = cached_model(("background_remover", _MODEL_NAME, str(device)), load_model)
|
|
22
|
+
|
|
23
|
+
# Data settings
|
|
24
|
+
transform_image = transforms.Compose(
|
|
25
|
+
[
|
|
26
|
+
transforms.Resize((1024, 1024)),
|
|
27
|
+
transforms.ToTensor(),
|
|
28
|
+
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
|
29
|
+
]
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
working_copy = image.copy()
|
|
33
|
+
input_images = transform_image(working_copy).unsqueeze(0).to(device)
|
|
34
|
+
|
|
35
|
+
# Prediction
|
|
36
|
+
with torch.no_grad():
|
|
37
|
+
preds = model(input_images)[-1].sigmoid().cpu()
|
|
38
|
+
pred = preds[0].squeeze()
|
|
39
|
+
pred_pil = transforms.ToPILImage()(pred)
|
|
40
|
+
mask = pred_pil.resize(working_copy.size)
|
|
41
|
+
working_copy.putalpha(mask)
|
|
42
|
+
|
|
43
|
+
return working_copy
|
dw/tasks/borders.py
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
from PIL import Image
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def add_border_and_mask(
|
|
5
|
+
image, zoom_all=1.0, zoom_left=0, zoom_right=0, zoom_up=0, zoom_down=0, overlap=0
|
|
6
|
+
):
|
|
7
|
+
"""Adds a black border around the image with individual side control and mask overlap"""
|
|
8
|
+
orig_width, orig_height = image.size
|
|
9
|
+
|
|
10
|
+
# Calculate padding for each side (in pixels)
|
|
11
|
+
left_pad = int(orig_width * zoom_left)
|
|
12
|
+
right_pad = int(orig_width * zoom_right)
|
|
13
|
+
top_pad = int(orig_height * zoom_up)
|
|
14
|
+
bottom_pad = int(orig_height * zoom_down)
|
|
15
|
+
|
|
16
|
+
# Calculate overlap in pixels
|
|
17
|
+
overlap_left = int(orig_width * overlap)
|
|
18
|
+
overlap_right = int(orig_width * overlap)
|
|
19
|
+
overlap_top = int(orig_height * overlap)
|
|
20
|
+
overlap_bottom = int(orig_height * overlap)
|
|
21
|
+
|
|
22
|
+
# If using the all-sides zoom, add it to each side
|
|
23
|
+
if zoom_all > 1.0:
|
|
24
|
+
extra_each_side = (zoom_all - 1.0) / 2
|
|
25
|
+
left_pad += int(orig_width * extra_each_side)
|
|
26
|
+
right_pad += int(orig_width * extra_each_side)
|
|
27
|
+
top_pad += int(orig_height * extra_each_side)
|
|
28
|
+
bottom_pad += int(orig_height * extra_each_side)
|
|
29
|
+
|
|
30
|
+
# Calculate new dimensions (ensure they're multiples of 32)
|
|
31
|
+
new_width = 32 * round((orig_width + left_pad + right_pad) / 32)
|
|
32
|
+
new_height = 32 * round((orig_height + top_pad + bottom_pad) / 32)
|
|
33
|
+
|
|
34
|
+
# Create new image with black border
|
|
35
|
+
bordered_image = Image.new("RGB", (new_width, new_height), (0, 0, 0))
|
|
36
|
+
# Paste original image in position
|
|
37
|
+
paste_x = left_pad
|
|
38
|
+
paste_y = top_pad
|
|
39
|
+
bordered_image.paste(image, (paste_x, paste_y))
|
|
40
|
+
|
|
41
|
+
# Create mask (white where the border is, black where the original image was)
|
|
42
|
+
mask = Image.new("L", (new_width, new_height), 255) # White background
|
|
43
|
+
# Paste black rectangle with overlap adjustment
|
|
44
|
+
mask.paste(
|
|
45
|
+
0,
|
|
46
|
+
(
|
|
47
|
+
paste_x + overlap_left, # Left edge moves right
|
|
48
|
+
paste_y + overlap_top, # Top edge moves down
|
|
49
|
+
paste_x + orig_width - overlap_right, # Right edge moves left
|
|
50
|
+
paste_y + orig_height - overlap_bottom, # Bottom edge moves up
|
|
51
|
+
),
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
return {"bordered_image": bordered_image, "mask": mask}
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def add_border_and_mask_with_size(image, width, height, overlap=0):
|
|
58
|
+
"""
|
|
59
|
+
Resizes the original image to fit within the target dimensions while maintaining
|
|
60
|
+
its aspect ratio, then adds borders as needed to reach the exact target size.
|
|
61
|
+
|
|
62
|
+
Args:
|
|
63
|
+
image: PIL Image object
|
|
64
|
+
width: Target width in pixels
|
|
65
|
+
height: Target height in pixels
|
|
66
|
+
overlap: Mask overlap parameter (0-1 range)
|
|
67
|
+
|
|
68
|
+
Returns:
|
|
69
|
+
Dictionary with 'bordered_image' and 'mask'
|
|
70
|
+
"""
|
|
71
|
+
# Ensure width and height are multiples of 32
|
|
72
|
+
width = 32 * round(width / 32)
|
|
73
|
+
height = 32 * round(height / 32)
|
|
74
|
+
|
|
75
|
+
# Get original dimensions
|
|
76
|
+
orig_width, orig_height = image.size
|
|
77
|
+
orig_aspect = orig_width / orig_height
|
|
78
|
+
target_aspect = width / height
|
|
79
|
+
|
|
80
|
+
# Resize image to fit within target dimensions while maintaining aspect ratio
|
|
81
|
+
if orig_aspect > target_aspect:
|
|
82
|
+
# Original is wider than target - fit width
|
|
83
|
+
new_width = width
|
|
84
|
+
new_height = int(width / orig_aspect)
|
|
85
|
+
resized_image = image.resize((new_width, new_height), Image.LANCZOS)
|
|
86
|
+
else:
|
|
87
|
+
# Original is taller than target - fit height
|
|
88
|
+
new_height = height
|
|
89
|
+
new_width = int(height * orig_aspect)
|
|
90
|
+
resized_image = image.resize((new_width, new_height), Image.LANCZOS)
|
|
91
|
+
|
|
92
|
+
# Now calculate padding to reach target dimensions
|
|
93
|
+
left_pad = (width - new_width) // 2
|
|
94
|
+
right_pad = width - new_width - left_pad
|
|
95
|
+
top_pad = (height - new_height) // 2
|
|
96
|
+
bottom_pad = height - new_height - top_pad
|
|
97
|
+
|
|
98
|
+
# Convert padding to zoom factors (relative to resized dimensions)
|
|
99
|
+
zoom_left = left_pad / new_width if new_width > 0 else 0
|
|
100
|
+
zoom_right = right_pad / new_width if new_width > 0 else 0
|
|
101
|
+
zoom_up = top_pad / new_height if new_height > 0 else 0
|
|
102
|
+
zoom_down = bottom_pad / new_height if new_height > 0 else 0
|
|
103
|
+
|
|
104
|
+
# Call the original function with calculated zoom parameters
|
|
105
|
+
return add_border_and_mask(
|
|
106
|
+
resized_image,
|
|
107
|
+
zoom_all=1.0,
|
|
108
|
+
zoom_left=zoom_left,
|
|
109
|
+
zoom_right=zoom_right,
|
|
110
|
+
zoom_up=zoom_up,
|
|
111
|
+
zoom_down=zoom_down,
|
|
112
|
+
overlap=overlap,
|
|
113
|
+
)
|
dw/tasks/compose_text.py
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
1
|
+
"""Assemble one block of text out of parts written once.
|
|
2
|
+
|
|
3
|
+
A multi-shot workflow says the same things about its characters in every
|
|
4
|
+
shot: who they are, what they are wearing, what their voice sounds like. The
|
|
5
|
+
engine has no interpolation - a reference is the whole value of an argument,
|
|
6
|
+
never a fragment spliced into one - which is deliberate (docs/WORKFLOW_GUIDE.md,
|
|
7
|
+
'no interpolation'): a `{{name}}` inside a prompt would make every prompt a
|
|
8
|
+
template language nobody declared, and resolution order for a half-substituted
|
|
9
|
+
string is a bottomless pit.
|
|
10
|
+
|
|
11
|
+
The way out is not interpolation but composition. A part is a whole value -
|
|
12
|
+
a variable, an earlier step's output, a stored prompt - and this task joins
|
|
13
|
+
parts in the order they are given. A character bible is then written once,
|
|
14
|
+
as a variable, and named by every shot that needs it; changing the voice
|
|
15
|
+
changes it everywhere, and nothing has to be hand-copied to stay in step.
|
|
16
|
+
|
|
17
|
+
Positional, not named: the parts are a list, joined in order. A named form
|
|
18
|
+
('{bible} says {line}') would be the interpolation the engine does not have,
|
|
19
|
+
one layer down.
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
import logging
|
|
23
|
+
|
|
24
|
+
logger = logging.getLogger("dw")
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def compose_text(parts, separator="\n\n", skip_empty=True):
|
|
28
|
+
"""Task command: join parts into one block of text.
|
|
29
|
+
|
|
30
|
+
Args:
|
|
31
|
+
parts: The parts to join, in order. Each is a whole value - usually
|
|
32
|
+
a "variable:", "prompt:" or "previous_result:" reference the
|
|
33
|
+
engine has already resolved. Numbers are written out; None is
|
|
34
|
+
dropped, so an optional part can be a variable left null
|
|
35
|
+
separator: What goes between the parts. Defaults to a blank line,
|
|
36
|
+
the paragraph break the prompt formats are written in
|
|
37
|
+
skip_empty: Drop parts that are None or empty. With it off, an
|
|
38
|
+
empty part still contributes its separator
|
|
39
|
+
|
|
40
|
+
Returns:
|
|
41
|
+
The joined text
|
|
42
|
+
|
|
43
|
+
Raises:
|
|
44
|
+
ValueError: If parts is not a list, or holds something that is not
|
|
45
|
+
text or a number - a dict or an image is a sign a reference
|
|
46
|
+
resolved to something other than what was meant
|
|
47
|
+
"""
|
|
48
|
+
if not isinstance(parts, list):
|
|
49
|
+
raise ValueError(
|
|
50
|
+
f"compose_text needs a list of parts to join, got {type(parts).__name__}"
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
pieces = []
|
|
54
|
+
for index, part in enumerate(parts):
|
|
55
|
+
if part is None:
|
|
56
|
+
if skip_empty:
|
|
57
|
+
continue
|
|
58
|
+
part = ""
|
|
59
|
+
if isinstance(part, bool) or not isinstance(part, (str, int, float)):
|
|
60
|
+
raise ValueError(
|
|
61
|
+
f"compose_text part {index} is a {type(part).__name__} - a part "
|
|
62
|
+
f"is text (or a number), and anything else means the reference "
|
|
63
|
+
f"in that position resolved to something other than text"
|
|
64
|
+
)
|
|
65
|
+
text = part if isinstance(part, str) else str(part)
|
|
66
|
+
if skip_empty and not text.strip():
|
|
67
|
+
continue
|
|
68
|
+
pieces.append(text)
|
|
69
|
+
|
|
70
|
+
if not pieces:
|
|
71
|
+
raise ValueError("compose_text was given nothing to join")
|
|
72
|
+
|
|
73
|
+
logger.debug(f"compose_text: joined {len(pieces)} of {len(parts)} parts")
|
|
74
|
+
return separator.join(pieces)
|
|
@@ -0,0 +1,300 @@
|
|
|
1
|
+
"""Concatenate videos - and the audio generated with them - into one video.
|
|
2
|
+
|
|
3
|
+
The standalone counterpart of what a chained pipeline step does internally:
|
|
4
|
+
frames are joined end to end with an optional head trim on every video after
|
|
5
|
+
the first, and audio tracks are joined at each seam with an equal-power
|
|
6
|
+
crossfade drawn from the trimmed-off material, so video and audio stay in
|
|
7
|
+
sync. Cuts have no trimmed material to fade with, so they can instead let the
|
|
8
|
+
outgoing tail ring on across the seam - see `audio_bleed_ms`.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
import logging
|
|
12
|
+
import os
|
|
13
|
+
|
|
14
|
+
from ..events import emit_log, emit_warning
|
|
15
|
+
from ..result import AudioVideo
|
|
16
|
+
from ..shots import measured_num_samples, nested_shots, shot_record, trimmed_shots
|
|
17
|
+
from .audio_utils import (
|
|
18
|
+
as_channels_samples,
|
|
19
|
+
bleed_join,
|
|
20
|
+
equal_power_crossfade_join,
|
|
21
|
+
fit_audio_to_frames,
|
|
22
|
+
frames_to_samples,
|
|
23
|
+
match_levels as match_track_levels,
|
|
24
|
+
resample_waveform,
|
|
25
|
+
warn_on_level_spread,
|
|
26
|
+
)
|
|
27
|
+
from .video_utils import check_same_frame_size, frames_as_pil_list, load_audio_video
|
|
28
|
+
|
|
29
|
+
logger = logging.getLogger("dw")
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def video_names(videos):
|
|
33
|
+
"""A name per video, for an error or a warning that has to say which one.
|
|
34
|
+
|
|
35
|
+
A caller passes a path, or a previous step's result; only the path says
|
|
36
|
+
anything by itself, so the rest are named by position - which is what a
|
|
37
|
+
six-entry `shots` list needs to be actionable ("24000 then 32000" does
|
|
38
|
+
not say which entry to fix). By the time this runs, an `asset:`/`output:`
|
|
39
|
+
reference has already been resolved to its absolute path on this server
|
|
40
|
+
(#390) - naming a shot by that path leaked server layout onto a consumer
|
|
41
|
+
surface, so a path is trimmed to its file name, the one part that means
|
|
42
|
+
anything off this box.
|
|
43
|
+
"""
|
|
44
|
+
return [
|
|
45
|
+
os.path.basename(original)
|
|
46
|
+
if isinstance(original, str)
|
|
47
|
+
else f"video {index + 1}"
|
|
48
|
+
for index, original in enumerate(videos)
|
|
49
|
+
]
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def concat_videos(
|
|
53
|
+
videos,
|
|
54
|
+
trim_frames=0,
|
|
55
|
+
crossfade_ms=75,
|
|
56
|
+
audio_bleed_ms=0,
|
|
57
|
+
audio_bleed_gain_db=0,
|
|
58
|
+
seam_fade_ms=None,
|
|
59
|
+
fps=None,
|
|
60
|
+
match_levels=None,
|
|
61
|
+
match_levels_dbfs=None,
|
|
62
|
+
sample_rate=None,
|
|
63
|
+
):
|
|
64
|
+
"""Concatenate a list of videos into a single AudioVideo.
|
|
65
|
+
|
|
66
|
+
Args:
|
|
67
|
+
videos: The videos to join, in order - frame lists, frame arrays,
|
|
68
|
+
AudioVideos (from previous_result references), or the path or URL
|
|
69
|
+
of a video file, which is read with the audio muxed into it. Give
|
|
70
|
+
each video its own entry: one previous_result reference naming a
|
|
71
|
+
step that produced several videos fans this step out over them,
|
|
72
|
+
one concatenation per video, rather than joining them
|
|
73
|
+
trim_frames: Frames dropped from the head of every video after the
|
|
74
|
+
first - the trim used when each video was generated from the
|
|
75
|
+
previous one's last frame
|
|
76
|
+
crossfade_ms: Equal-power crossfade at each audio seam, drawn from
|
|
77
|
+
the trimmed material - so it has no effect when trim_frames is 0,
|
|
78
|
+
which is where every cut-based workflow sits. A hard cut's seam is
|
|
79
|
+
shaped by audio_bleed_ms or seam_fade_ms instead
|
|
80
|
+
audio_bleed_ms: How long the outgoing video's tail rings on over the head
|
|
81
|
+
of the next one, at seams with nothing trimmed to crossfade. For
|
|
82
|
+
cut-based workflows, where every shot is generated independently and
|
|
83
|
+
a running laugh track would otherwise butt-join into silence.
|
|
84
|
+
0 (the default) leaves the seam as a plain declicked join
|
|
85
|
+
audio_bleed_gain_db: Gain applied to the bled tail before it is added,
|
|
86
|
+
in dB. 0 (the default) is unchanged, full-scale, matching the
|
|
87
|
+
outgoing material exactly; negative ducks a tail that would
|
|
88
|
+
otherwise push the seam over 0 dBFS, or that reads as too present
|
|
89
|
+
against the incoming shot. Has no effect when audio_bleed_ms is 0
|
|
90
|
+
seam_fade_ms: Fade applied on each side of a seam that gets neither a
|
|
91
|
+
crossfade nor a bleed. Defaults to the few milliseconds that keep a
|
|
92
|
+
butt-join from clicking; raise it to a hundred or so for a graceful
|
|
93
|
+
hard cut on tonal material, which a bleed would only stutter. It is
|
|
94
|
+
the wrong tool for a continuous bed such as a laugh track or room
|
|
95
|
+
tone: a fade only deepens the hole a bleed is there to cover
|
|
96
|
+
fps: Frame rate of the videos - required to join audio when
|
|
97
|
+
trimming, and the rate the joined file is written at unless
|
|
98
|
+
the step's result.fps overrides it
|
|
99
|
+
match_levels: Even the shots' loudness out before joining -
|
|
100
|
+
"rms" matches perceived level (the measurement
|
|
101
|
+
get_gallery_metadata reports as mean_dbfs), "peak" matches the
|
|
102
|
+
loudest sample. Off by default. Shots generated independently
|
|
103
|
+
land 10 dB apart routinely, and that jump is the one seam
|
|
104
|
+
artifact none of the fade controls can hide, because it is not
|
|
105
|
+
at the seam but either side of it. Left off, a spread wide
|
|
106
|
+
enough to hear is logged as a warning
|
|
107
|
+
match_levels_dbfs: The level match_levels moves every shot to -
|
|
108
|
+
defaults to -1 dBFS for "peak" and -20 dBFS for "rms". A shot
|
|
109
|
+
that would clip at the target is held at -0.5 dBFS peak instead,
|
|
110
|
+
reported as a match_levels_held warning, with a per-shot log
|
|
111
|
+
event naming the hold
|
|
112
|
+
sample_rate: The rate the joined soundtrack is at. Shots that come
|
|
113
|
+
from different sources routinely carry different rates - a 24 kHz
|
|
114
|
+
voice clip paired onto a 32 kHz generation - and unlike a level
|
|
115
|
+
jump that difference has no editorial meaning, so by default the
|
|
116
|
+
highest rate among the inputs is chosen and the rest are
|
|
117
|
+
resampled up to it, with a warning naming which. Give this to pin
|
|
118
|
+
the target instead (#108)
|
|
119
|
+
|
|
120
|
+
Returns:
|
|
121
|
+
One AudioVideo; its audio is None when no input video carries any
|
|
122
|
+
"""
|
|
123
|
+
if not isinstance(videos, list) or not videos:
|
|
124
|
+
raise ValueError("concat_videos needs a non-empty list of videos")
|
|
125
|
+
|
|
126
|
+
# Named before they are loaded: a path is the only thing that names
|
|
127
|
+
# itself, and the load below replaces it with what it holds
|
|
128
|
+
names = video_names(videos)
|
|
129
|
+
# A shot an earlier run already wrote is loaded here rather than by
|
|
130
|
+
# gather_videos, which reads frames only and would join it silent
|
|
131
|
+
videos = [load_audio_video(v) if isinstance(v, str) else v for v in videos]
|
|
132
|
+
clips = [frames_as_pil_list(v) for v in videos]
|
|
133
|
+
check_same_frame_size(clips, "concat_videos")
|
|
134
|
+
|
|
135
|
+
# Every track up front rather than one at a time: levels are matched
|
|
136
|
+
# across the whole set, so the last shot's loudness has to be known
|
|
137
|
+
# before the first one is scaled
|
|
138
|
+
waveforms = [
|
|
139
|
+
(
|
|
140
|
+
as_channels_samples(video.audio)
|
|
141
|
+
if isinstance(video, AudioVideo) and video.audio is not None
|
|
142
|
+
else None
|
|
143
|
+
)
|
|
144
|
+
for video in videos
|
|
145
|
+
]
|
|
146
|
+
# One rate before anything is joined. Shots assembled from different
|
|
147
|
+
# sources disagree routinely, and the disagreement carries no meaning -
|
|
148
|
+
# so it is converted rather than refused, which is what made an agent
|
|
149
|
+
# invent a resample_audio step by hand (#108)
|
|
150
|
+
rates = [
|
|
151
|
+
video.sample_rate
|
|
152
|
+
for video, waveform in zip(videos, waveforms)
|
|
153
|
+
if waveform is not None and video.sample_rate
|
|
154
|
+
]
|
|
155
|
+
sample_rate = sample_rate or (max(rates) if rates else None)
|
|
156
|
+
if rates and len(set(rates)) == 1 and rates[0] != sample_rate:
|
|
157
|
+
# The inputs agree and the caller pinned another rate: converting
|
|
158
|
+
# to what was asked for is not a decision made on its behalf (#453)
|
|
159
|
+
emit_log(
|
|
160
|
+
f"concat_videos: resampling every track from {rates[0]} Hz to the "
|
|
161
|
+
f"requested {sample_rate} Hz",
|
|
162
|
+
command="concat_videos",
|
|
163
|
+
sample_rate=sample_rate,
|
|
164
|
+
)
|
|
165
|
+
elif rates and any(rate != sample_rate for rate in rates):
|
|
166
|
+
# emit_warning rather than logger.warning, for the reason the level
|
|
167
|
+
# spread below is emitted: resampling every track is an audio
|
|
168
|
+
# decision made on the caller's behalf, and a caller reading the job
|
|
169
|
+
# over the API or MCP sees the warnings list and nothing else - the
|
|
170
|
+
# conversion landing silently is worse than the loud failure it
|
|
171
|
+
# replaced (#108)
|
|
172
|
+
per_video = {
|
|
173
|
+
name: video.sample_rate
|
|
174
|
+
for name, video in zip(names, videos)
|
|
175
|
+
if isinstance(video, AudioVideo) and video.audio is not None
|
|
176
|
+
}
|
|
177
|
+
emit_warning(
|
|
178
|
+
"concat_videos: videos carry audio at different sample rates ("
|
|
179
|
+
+ ", ".join(f"{name}: {rate} Hz" for name, rate in per_video.items())
|
|
180
|
+
+ f") - resampling them all to {sample_rate} Hz. Pass "
|
|
181
|
+
"'sample_rate' to pin a different target, or resample ahead of "
|
|
182
|
+
"this step with the 'resample_audio' task.",
|
|
183
|
+
kind="sample_rate_mismatch",
|
|
184
|
+
command="concat_videos",
|
|
185
|
+
sample_rate=sample_rate,
|
|
186
|
+
sample_rates=per_video,
|
|
187
|
+
)
|
|
188
|
+
waveforms = [
|
|
189
|
+
(
|
|
190
|
+
waveform
|
|
191
|
+
if waveform is None
|
|
192
|
+
or not video.sample_rate
|
|
193
|
+
or video.sample_rate == sample_rate
|
|
194
|
+
else resample_waveform(waveform, video.sample_rate, sample_rate)
|
|
195
|
+
)
|
|
196
|
+
for video, waveform in zip(videos, waveforms)
|
|
197
|
+
]
|
|
198
|
+
|
|
199
|
+
if match_levels:
|
|
200
|
+
waveforms = match_track_levels(waveforms, match_levels, match_levels_dbfs)
|
|
201
|
+
else:
|
|
202
|
+
warn_on_level_spread(waveforms)
|
|
203
|
+
|
|
204
|
+
frames = []
|
|
205
|
+
audio = None
|
|
206
|
+
audio_native_rate = None
|
|
207
|
+
# Where each video landed, measured on the joined picture and track as
|
|
208
|
+
# they grow - never derived from the frame count, so a track that runs
|
|
209
|
+
# long shows up here as the samples it actually took (#378)
|
|
210
|
+
shots = []
|
|
211
|
+
|
|
212
|
+
for index, (video, clip) in enumerate(zip(videos, clips)):
|
|
213
|
+
head_trim = trim_frames if index > 0 else 0
|
|
214
|
+
start_frame = len(frames)
|
|
215
|
+
start_sample = audio.shape[1] if audio is not None else 0
|
|
216
|
+
frames.extend(clip[head_trim:])
|
|
217
|
+
inner = getattr(video, "shots", None)
|
|
218
|
+
if inner:
|
|
219
|
+
video_shots = nested_shots(
|
|
220
|
+
trimmed_shots(inner, head_trim),
|
|
221
|
+
start_frame,
|
|
222
|
+
start_sample if waveforms[index] is not None else None,
|
|
223
|
+
getattr(video, "sample_rate", None),
|
|
224
|
+
sample_rate,
|
|
225
|
+
)
|
|
226
|
+
else:
|
|
227
|
+
video_shots = [
|
|
228
|
+
shot_record(
|
|
229
|
+
names[index], start_frame, len(frames) - start_frame, start_sample
|
|
230
|
+
)
|
|
231
|
+
]
|
|
232
|
+
# Which input this shot came from - named_shots (dw/shots.py) uses
|
|
233
|
+
# it to place a step's override name on the right shot once an
|
|
234
|
+
# earlier input has nested more than one of its own (#432)
|
|
235
|
+
for shot in video_shots:
|
|
236
|
+
shot["source_index"] = index
|
|
237
|
+
shots.extend(video_shots)
|
|
238
|
+
|
|
239
|
+
if waveforms[index] is None:
|
|
240
|
+
continue
|
|
241
|
+
|
|
242
|
+
waveform = waveforms[index]
|
|
243
|
+
if audio is None:
|
|
244
|
+
audio = waveform
|
|
245
|
+
audio_native_rate = video.sample_rate
|
|
246
|
+
continue
|
|
247
|
+
|
|
248
|
+
if head_trim > 0 and fps is None:
|
|
249
|
+
raise ValueError(
|
|
250
|
+
"concat_videos needs 'fps' to trim audio in step with the frames"
|
|
251
|
+
)
|
|
252
|
+
|
|
253
|
+
trim_samples = (
|
|
254
|
+
frames_to_samples(head_trim, fps, sample_rate) if head_trim else 0
|
|
255
|
+
)
|
|
256
|
+
if trim_samples == 0 and audio_bleed_ms:
|
|
257
|
+
audio = bleed_join(
|
|
258
|
+
audio,
|
|
259
|
+
waveform,
|
|
260
|
+
sample_rate,
|
|
261
|
+
audio_bleed_ms,
|
|
262
|
+
seam_fade_ms,
|
|
263
|
+
audio_bleed_gain_db,
|
|
264
|
+
native_sample_rate=audio_native_rate,
|
|
265
|
+
)
|
|
266
|
+
else:
|
|
267
|
+
audio = equal_power_crossfade_join(
|
|
268
|
+
audio,
|
|
269
|
+
waveform[:, :trim_samples],
|
|
270
|
+
waveform[:, trim_samples:],
|
|
271
|
+
sample_rate,
|
|
272
|
+
crossfade_ms,
|
|
273
|
+
seam_fade_ms,
|
|
274
|
+
)
|
|
275
|
+
audio_native_rate = video.sample_rate
|
|
276
|
+
|
|
277
|
+
# The rate the caller declared, else the rate the first input carries -
|
|
278
|
+
# either beats the result's 8 fps default (#84)
|
|
279
|
+
written_fps = fps or next(
|
|
280
|
+
(v.fps for v in videos if getattr(v, "fps", None)),
|
|
281
|
+
None,
|
|
282
|
+
)
|
|
283
|
+
# Reconciled against the frame grid before shots are measured (#435), so
|
|
284
|
+
# an input already short of its own grid does not carry its shortfall
|
|
285
|
+
# into this join's shot map and compound in a later one
|
|
286
|
+
audio = fit_audio_to_frames(
|
|
287
|
+
audio, sample_rate, len(frames), written_fps, "concat_videos"
|
|
288
|
+
)
|
|
289
|
+
|
|
290
|
+
# A seam's crossfade leaves the samples before it where they were, so a
|
|
291
|
+
# shot's track is everything up to where the next measured one began
|
|
292
|
+
measured_num_samples(shots, _length(audio) if audio is not None else None)
|
|
293
|
+
|
|
294
|
+
logger.debug(f"Concatenated {len(videos)} videos into {len(frames)} frames")
|
|
295
|
+
return AudioVideo(frames, audio, sample_rate, fps=written_fps, shots=shots)
|
|
296
|
+
|
|
297
|
+
|
|
298
|
+
def _length(audio):
|
|
299
|
+
"""How many samples a joined track holds, 0 for none."""
|
|
300
|
+
return 0 if audio is None else audio.shape[1]
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import numpy as np
|
|
3
|
+
from transformers import pipeline
|
|
4
|
+
from torchvision import transforms
|
|
5
|
+
|
|
6
|
+
from .. import preferred_task_dtype
|
|
7
|
+
from .model_cache import cached_model
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def make_hint_tensor(image, device, dtype=None):
|
|
11
|
+
"""Estimate depth and return it as a hint tensor for a controlnet pipeline.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
image: Image to estimate depth from
|
|
15
|
+
device: Device to run the estimator on and place the hint on
|
|
16
|
+
dtype: Dtype of the hint, defaulting to the one the device works best in.
|
|
17
|
+
The hint has to match the dtype of the pipeline that consumes it.
|
|
18
|
+
|
|
19
|
+
Returns:
|
|
20
|
+
Depth hint as a tensor of shape (1, 3, height, width)
|
|
21
|
+
"""
|
|
22
|
+
depth_estimator = cached_model(
|
|
23
|
+
("depth_estimator", str(device)),
|
|
24
|
+
lambda: pipeline("depth-estimation", device=device),
|
|
25
|
+
)
|
|
26
|
+
|
|
27
|
+
image = depth_estimator(image)["depth"]
|
|
28
|
+
image = np.array(image)
|
|
29
|
+
image = image[:, :, None]
|
|
30
|
+
image = np.concatenate([image, image, image], axis=2)
|
|
31
|
+
detected_map = torch.from_numpy(image).float() / 255.0
|
|
32
|
+
hint = detected_map.permute(2, 0, 1)
|
|
33
|
+
|
|
34
|
+
if dtype is None:
|
|
35
|
+
dtype = preferred_task_dtype(device)
|
|
36
|
+
|
|
37
|
+
return hint.unsqueeze(0).to(device=device, dtype=dtype)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def make_hint_image(image, device, dtype=None):
|
|
41
|
+
"""Estimate depth and return it as an image.
|
|
42
|
+
|
|
43
|
+
Args:
|
|
44
|
+
image: Image to estimate depth from
|
|
45
|
+
device: Device to run the estimator on
|
|
46
|
+
dtype: Dtype to compute the hint in - see make_hint_tensor
|
|
47
|
+
|
|
48
|
+
Returns:
|
|
49
|
+
Depth map as a PIL image
|
|
50
|
+
"""
|
|
51
|
+
hint = make_hint_tensor(image, device, dtype)
|
|
52
|
+
# Convert the tensor to a Pillow image
|
|
53
|
+
to_pil = transforms.ToPILImage()
|
|
54
|
+
return to_pil(hint[0].float().cpu())
|
|
@@ -0,0 +1,109 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Diffusion-based image upscaling via Stable Diffusion upscale pipelines.
|
|
3
|
+
|
|
4
|
+
Provides text-guided upscaling with better detail recovery than
|
|
5
|
+
traditional super-resolution models, especially for faces and textures.
|
|
6
|
+
|
|
7
|
+
Supports two modes:
|
|
8
|
+
- "x4" (default): StableDiffusionUpscalePipeline (4x, stabilityai/stable-diffusion-x4-upscaler)
|
|
9
|
+
- "x2": StableDiffusionLatentUpscalePipeline (2x, stabilityai/sd-x2-latent-upscaler)
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
import logging
|
|
13
|
+
import torch
|
|
14
|
+
import diffusers
|
|
15
|
+
from .. import preferred_task_dtype
|
|
16
|
+
from .model_cache import cached_model
|
|
17
|
+
|
|
18
|
+
logger = logging.getLogger("dw")
|
|
19
|
+
|
|
20
|
+
_MODELS = {
|
|
21
|
+
"x4": {
|
|
22
|
+
"pipeline_class": "StableDiffusionUpscalePipeline",
|
|
23
|
+
"model_name": "stabilityai/stable-diffusion-x4-upscaler",
|
|
24
|
+
},
|
|
25
|
+
"x2": {
|
|
26
|
+
"pipeline_class": "StableDiffusionLatentUpscalePipeline",
|
|
27
|
+
"model_name": "stabilityai/sd-x2-latent-upscaler",
|
|
28
|
+
},
|
|
29
|
+
}
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def diffusion_upscale(image, device="cpu", **kwargs):
|
|
33
|
+
"""Upscale an image using a Stable Diffusion upscale pipeline.
|
|
34
|
+
|
|
35
|
+
Args:
|
|
36
|
+
image: PIL Image to upscale.
|
|
37
|
+
device: Target device ("cuda", "mps", "cpu").
|
|
38
|
+
**kwargs:
|
|
39
|
+
prompt: Text guidance for upscaling (default: "").
|
|
40
|
+
negative_prompt: Negative text guidance (default: None).
|
|
41
|
+
mode: "x4" or "x2" (default: "x4").
|
|
42
|
+
model_name: Override the default model for the selected mode.
|
|
43
|
+
num_inference_steps: Denoising steps (default: 25).
|
|
44
|
+
guidance_scale: Classifier-free guidance scale (default: 9.0).
|
|
45
|
+
noise_level: Noise level for x4 mode (default: 20, ignored for x2).
|
|
46
|
+
|
|
47
|
+
Returns:
|
|
48
|
+
PIL Image (upscaled).
|
|
49
|
+
"""
|
|
50
|
+
mode = kwargs.get("mode", "x4")
|
|
51
|
+
if mode not in _MODELS:
|
|
52
|
+
raise ValueError(f"mode must be one of {sorted(_MODELS.keys())}, got '{mode}'")
|
|
53
|
+
|
|
54
|
+
config = _MODELS[mode]
|
|
55
|
+
model_name = kwargs.get("model_name", config["model_name"])
|
|
56
|
+
prompt = kwargs.get("prompt", "")
|
|
57
|
+
negative_prompt = kwargs.get("negative_prompt", None)
|
|
58
|
+
num_inference_steps = int(kwargs.get("num_inference_steps", 25))
|
|
59
|
+
guidance_scale = float(kwargs.get("guidance_scale", 9.0))
|
|
60
|
+
noise_level = int(kwargs.get("noise_level", 20))
|
|
61
|
+
|
|
62
|
+
pipeline_class = getattr(diffusers, config["pipeline_class"])
|
|
63
|
+
|
|
64
|
+
dtype = preferred_task_dtype(device)
|
|
65
|
+
|
|
66
|
+
def load_pipe():
|
|
67
|
+
logger.info(f"Loading {config['pipeline_class']} from {model_name} to {device}")
|
|
68
|
+
pipe = pipeline_class.from_pretrained(
|
|
69
|
+
model_name,
|
|
70
|
+
torch_dtype=dtype,
|
|
71
|
+
)
|
|
72
|
+
pipe.to(device)
|
|
73
|
+
return pipe
|
|
74
|
+
|
|
75
|
+
pipe = cached_model(
|
|
76
|
+
(
|
|
77
|
+
"diffusion_upscale",
|
|
78
|
+
config["pipeline_class"],
|
|
79
|
+
model_name,
|
|
80
|
+
str(device),
|
|
81
|
+
str(dtype),
|
|
82
|
+
),
|
|
83
|
+
load_pipe,
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
call_kwargs = {
|
|
87
|
+
"prompt": prompt,
|
|
88
|
+
"image": image,
|
|
89
|
+
"num_inference_steps": num_inference_steps,
|
|
90
|
+
"guidance_scale": guidance_scale,
|
|
91
|
+
}
|
|
92
|
+
|
|
93
|
+
if negative_prompt is not None:
|
|
94
|
+
call_kwargs["negative_prompt"] = negative_prompt
|
|
95
|
+
|
|
96
|
+
if mode == "x4":
|
|
97
|
+
call_kwargs["noise_level"] = noise_level
|
|
98
|
+
|
|
99
|
+
logger.info(
|
|
100
|
+
f"Upscaling {image.width}x{image.height} with {mode} mode, "
|
|
101
|
+
f"{num_inference_steps} steps"
|
|
102
|
+
)
|
|
103
|
+
|
|
104
|
+
with torch.inference_mode():
|
|
105
|
+
result = pipe(**call_kwargs)
|
|
106
|
+
|
|
107
|
+
output = result.images[0]
|
|
108
|
+
logger.info(f"Upscaled to {output.width}x{output.height}")
|
|
109
|
+
return output
|