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/video_utils.py
ADDED
|
@@ -0,0 +1,624 @@
|
|
|
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 logging
|
|
11
|
+
import math
|
|
12
|
+
import re
|
|
13
|
+
|
|
14
|
+
import numpy
|
|
15
|
+
import torch
|
|
16
|
+
from PIL import Image
|
|
17
|
+
|
|
18
|
+
from ..result import AudioVideo
|
|
19
|
+
|
|
20
|
+
logger = logging.getLogger("dw")
|
|
21
|
+
|
|
22
|
+
# A location that names a scheme is a URL, whatever the scheme
|
|
23
|
+
_URL_SCHEME = re.compile(r"^[a-zA-Z][a-zA-Z0-9+.\-]*://")
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def process_video(video, processor, device, kwargs):
|
|
27
|
+
processor = processor.lower()
|
|
28
|
+
|
|
29
|
+
if processor == "get_frame":
|
|
30
|
+
return get_frame(video, kwargs.get("frame_index", 0))
|
|
31
|
+
|
|
32
|
+
if processor == "get_last_frame":
|
|
33
|
+
return get_frame(video, -1)
|
|
34
|
+
|
|
35
|
+
if processor == "get_first_frame":
|
|
36
|
+
return get_frame(video, 0)
|
|
37
|
+
|
|
38
|
+
raise Exception(f"Unknown video processor type: {processor}")
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class VideoFileReference:
|
|
42
|
+
"""A 'video' argument realized to a file on disk rather than an in-memory
|
|
43
|
+
clip - built by dw/arguments.py's _realize_lazy_frame_arguments so
|
|
44
|
+
get_frame can seek to the one frame it needs instead of decoding the
|
|
45
|
+
whole file (#367), and so an assessment probe streams the file, soundtrack
|
|
46
|
+
and all (#387). Not a public shape; nothing else constructs or consumes
|
|
47
|
+
one."""
|
|
48
|
+
|
|
49
|
+
__slots__ = ("path",)
|
|
50
|
+
|
|
51
|
+
def __init__(self, path):
|
|
52
|
+
self.path = path
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def get_frame(video, frame_index=0):
|
|
56
|
+
"""Pull one frame out of a video as a PIL image.
|
|
57
|
+
|
|
58
|
+
Args:
|
|
59
|
+
video: List of PIL images, numpy array or torch tensor of frames, an
|
|
60
|
+
AudioVideo, a one-video batch wrapping any of those, or a
|
|
61
|
+
VideoFileReference naming a file this call reads by seeking
|
|
62
|
+
rather than decoding in full
|
|
63
|
+
frame_index: Frame to extract, 0-based; negative indexes count from
|
|
64
|
+
the end (-1 is the last frame). Past either end of the clip
|
|
65
|
+
raises an error naming the clip's frame count
|
|
66
|
+
|
|
67
|
+
Returns:
|
|
68
|
+
The frame as a PIL image
|
|
69
|
+
"""
|
|
70
|
+
if isinstance(video, VideoFileReference):
|
|
71
|
+
from ..media_frames import frames_at
|
|
72
|
+
|
|
73
|
+
return frames_at(video.path, [f"frame:{frame_index}"])[0]["image"]
|
|
74
|
+
return extract_frame(video, frame_index)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def extract_frame(video, index):
|
|
78
|
+
"""Pull one frame out of a video, whatever the video's in-memory shape.
|
|
79
|
+
|
|
80
|
+
Args:
|
|
81
|
+
video: List of PIL images, numpy array or torch tensor of frames,
|
|
82
|
+
an AudioVideo, or a one-video batch wrapping any of those
|
|
83
|
+
index: Frame to extract; negative indexes count from the end
|
|
84
|
+
|
|
85
|
+
Returns:
|
|
86
|
+
The frame as a PIL image. Frames that already are PIL images are
|
|
87
|
+
returned as-is, not copied.
|
|
88
|
+
"""
|
|
89
|
+
return _to_pil(_frames_of(video)[index])
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def frame_count(video):
|
|
93
|
+
"""Number of frames in a video of any supported shape."""
|
|
94
|
+
return len(_frames_of(video))
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def check_same_frame_size(clips, task_name):
|
|
98
|
+
"""Refuse to join clips whose frames disagree in size.
|
|
99
|
+
|
|
100
|
+
Args:
|
|
101
|
+
clips: The clips about to be joined, each a PIL frame list or a
|
|
102
|
+
(frames, height, width, channels) array
|
|
103
|
+
task_name: Named in the error
|
|
104
|
+
|
|
105
|
+
Joining is frame-by-frame concatenation, which either fails deep in numpy
|
|
106
|
+
or, for a PIL list, produces a film that changes size mid-cut. Naming the
|
|
107
|
+
two sizes points at the shot that was rendered differently rather than at
|
|
108
|
+
the join.
|
|
109
|
+
"""
|
|
110
|
+
sizes = []
|
|
111
|
+
for clip in clips:
|
|
112
|
+
if isinstance(clip, numpy.ndarray):
|
|
113
|
+
sizes.append((int(clip.shape[2]), int(clip.shape[1])))
|
|
114
|
+
else:
|
|
115
|
+
sizes.append(tuple(clip[0].size) if len(clip) else None)
|
|
116
|
+
first = next((size for size in sizes if size is not None), None)
|
|
117
|
+
for index, size in enumerate(sizes):
|
|
118
|
+
if size is not None and size != first:
|
|
119
|
+
raise ValueError(
|
|
120
|
+
f"{task_name} needs every video at one size: video 0 is "
|
|
121
|
+
f"{first[0]}x{first[1]}, video {index} is {size[0]}x{size[1]}"
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def frames_as_pil_list(video):
|
|
126
|
+
"""The video's frames as a list of PIL images.
|
|
127
|
+
|
|
128
|
+
Frames that already are PIL images are carried over by identity; array and
|
|
129
|
+
tensor frames are converted the way extract_frame converts them.
|
|
130
|
+
"""
|
|
131
|
+
return [_to_pil(frame) for frame in _frames_of(video)]
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
def frames_as_array(video):
|
|
135
|
+
"""The video's frames as one (frames, height, width, channels) uint8 array.
|
|
136
|
+
|
|
137
|
+
The shape an argument that takes frames rather than a video wants - LTX-2's
|
|
138
|
+
keyframe conditions and IC-LoRA references, which the workflow hands what an
|
|
139
|
+
earlier step generated. One array is also one artifact, where a list of frames
|
|
140
|
+
would become one artifact per frame and multiply the step that consumed it.
|
|
141
|
+
|
|
142
|
+
Frames that already are a channels-last RGB array are converted in a single
|
|
143
|
+
operation; anything else goes through the same per-frame conversion
|
|
144
|
+
extract_frame uses.
|
|
145
|
+
"""
|
|
146
|
+
frames = _frames_of(video)
|
|
147
|
+
|
|
148
|
+
if isinstance(frames, numpy.ndarray) and frames.ndim == 4 and frames.shape[-1] == 3:
|
|
149
|
+
if frames.dtype == numpy.uint8:
|
|
150
|
+
return frames
|
|
151
|
+
# Float frames are [0, 1] - diffusers' np output convention
|
|
152
|
+
return (numpy.clip(frames, 0.0, 1.0) * 255).round().astype(numpy.uint8)
|
|
153
|
+
|
|
154
|
+
return numpy.stack(
|
|
155
|
+
[numpy.asarray(_to_pil(frame).convert("RGB")) for frame in frames]
|
|
156
|
+
)
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def loop_frames(video, num_frames):
|
|
160
|
+
"""Task command: a run of exactly `num_frames` frames, made by repeating
|
|
161
|
+
what it is given.
|
|
162
|
+
|
|
163
|
+
The video analogue of `loop_audio`, and it exists for the same reason: a
|
|
164
|
+
conditioning input has a length its model was trained to read, and the
|
|
165
|
+
material to hand is usually shorter. LTX-2.5's Ingredients IC-LoRA is
|
|
166
|
+
the live case - the reference is a single still sheet, and the model
|
|
167
|
+
wants it as a static video of at least 121 frames at the output's own
|
|
168
|
+
length, because a shorter reference misses the 121-frame read bucket it
|
|
169
|
+
was trained on.
|
|
170
|
+
|
|
171
|
+
A still is repeated; a run of frames laps round from its first frame.
|
|
172
|
+
No crossfade, unlike the audio version: these frames are read as
|
|
173
|
+
reference latents rather than watched, so a visible cut at the lap is
|
|
174
|
+
not a defect and blending two frames of a reference sheet would be.
|
|
175
|
+
|
|
176
|
+
Args:
|
|
177
|
+
video: Frames in any shape a result carries, or a still image - but
|
|
178
|
+
the `video` argument loads *video files* by convention (#347), so
|
|
179
|
+
a still on disk has to be passed as
|
|
180
|
+
`{"media_type": "image", "location": "asset:x.png"}` rather than
|
|
181
|
+
a bare path or `asset:`/`output:` reference; a still made earlier
|
|
182
|
+
in the same workflow is `previous_result:<image step>`
|
|
183
|
+
num_frames: How many frames to hand back, one or more
|
|
184
|
+
|
|
185
|
+
Returns:
|
|
186
|
+
A (num_frames, height, width, channels) float32 array scaled to
|
|
187
|
+
[0, 1] - diffusers' own np frame convention, and what
|
|
188
|
+
`LTX2ReferenceCondition.frames` and its kin need: a raw ndarray
|
|
189
|
+
reaches `VaeImageProcessor.preprocess` untouched, with no /255
|
|
190
|
+
rescaling applied along the way, so a uint8 [0, 255] array read as
|
|
191
|
+
already-scaled data is 255x too bright (#444). Not for a
|
|
192
|
+
keyframe (`LTX2VideoCondition`): its ndarray path expects uint8
|
|
193
|
+
[0, 255] and refuses a float frame when `crf` is set -
|
|
194
|
+
`frames_as_array` is the shape for that
|
|
195
|
+
"""
|
|
196
|
+
if isinstance(num_frames, str):
|
|
197
|
+
try:
|
|
198
|
+
num_frames = int(num_frames)
|
|
199
|
+
except ValueError:
|
|
200
|
+
raise ValueError(
|
|
201
|
+
f"loop_frames needs 'num_frames' as a whole number, got {num_frames!r}"
|
|
202
|
+
)
|
|
203
|
+
if not isinstance(num_frames, int) or isinstance(num_frames, bool):
|
|
204
|
+
raise ValueError(
|
|
205
|
+
f"loop_frames needs 'num_frames' as a whole number, got {num_frames!r}"
|
|
206
|
+
)
|
|
207
|
+
if num_frames < 1:
|
|
208
|
+
raise ValueError(
|
|
209
|
+
f"loop_frames needs 'num_frames' of at least 1, got {num_frames}"
|
|
210
|
+
)
|
|
211
|
+
|
|
212
|
+
# A lone still is the Ingredients case, and `_frames_of` does not take
|
|
213
|
+
# one - a reference sheet is an image, not a one-frame video
|
|
214
|
+
frames = frames_as_array([video] if _is_frame(video) else video)
|
|
215
|
+
if len(frames) == 0:
|
|
216
|
+
raise ValueError("loop_frames was given no frames to repeat")
|
|
217
|
+
laps = -(-num_frames // len(frames)) # ceiling, so the last lap is trimmed
|
|
218
|
+
looped = numpy.concatenate([frames] * laps, axis=0)[:num_frames]
|
|
219
|
+
return (looped.astype(numpy.float32) / 255.0).clip(0.0, 1.0)
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def frame_grid(video, count=12, columns=None, tile_width=320, label=True):
|
|
223
|
+
"""Task command: tile evenly sampled frames of a video into one contact
|
|
224
|
+
sheet - a preview of a clip's shape without authoring a frames-extraction
|
|
225
|
+
workflow (#245).
|
|
226
|
+
|
|
227
|
+
Args:
|
|
228
|
+
video: Frames in any shape a result carries
|
|
229
|
+
count: How many frames to sample, spaced evenly across the clip's
|
|
230
|
+
full duration (including its first and last frame). Clamped to
|
|
231
|
+
the clip's own frame count when the clip is shorter
|
|
232
|
+
columns: Tiles per row. Defaults to a grid biased wide - clips are
|
|
233
|
+
usually landscape - with the last row left-justified when
|
|
234
|
+
`count` is not a perfect multiple of it
|
|
235
|
+
tile_width: Width in pixels of each tile; height follows the source
|
|
236
|
+
frame's aspect ratio
|
|
237
|
+
label: Burn the sampled timestamp (or frame index, when the video
|
|
238
|
+
carries no frame rate) into each tile's corner
|
|
239
|
+
|
|
240
|
+
Returns:
|
|
241
|
+
One PIL image, the tiled contact sheet
|
|
242
|
+
"""
|
|
243
|
+
count = _positive_int(count, "frame_grid", "count")
|
|
244
|
+
if columns is not None:
|
|
245
|
+
columns = _positive_int(columns, "frame_grid", "columns")
|
|
246
|
+
tile_width = _positive_int(tile_width, "frame_grid", "tile_width")
|
|
247
|
+
if not isinstance(label, bool):
|
|
248
|
+
raise ValueError(f"frame_grid needs 'label' as true or false, got {label!r}")
|
|
249
|
+
|
|
250
|
+
total = frame_count(video)
|
|
251
|
+
if total == 0:
|
|
252
|
+
raise ValueError("frame_grid was given a video with no frames")
|
|
253
|
+
count = min(count, total)
|
|
254
|
+
fps = getattr(video, "fps", None)
|
|
255
|
+
|
|
256
|
+
indices = _evenly_spaced_indices(total, count)
|
|
257
|
+
tiles = [
|
|
258
|
+
_grid_tile(extract_frame(video, index), index, fps, tile_width, label)
|
|
259
|
+
for index in indices
|
|
260
|
+
]
|
|
261
|
+
|
|
262
|
+
if columns is None:
|
|
263
|
+
columns = _default_columns(len(tiles))
|
|
264
|
+
return _compose_grid(tiles, columns)
|
|
265
|
+
|
|
266
|
+
|
|
267
|
+
def _positive_int(value, command, name):
|
|
268
|
+
if isinstance(value, str):
|
|
269
|
+
try:
|
|
270
|
+
value = int(value)
|
|
271
|
+
except ValueError:
|
|
272
|
+
raise ValueError(
|
|
273
|
+
f"{command} needs '{name}' as a whole number, got {value!r}"
|
|
274
|
+
)
|
|
275
|
+
if not isinstance(value, int) or isinstance(value, bool):
|
|
276
|
+
raise ValueError(f"{command} needs '{name}' as a whole number, got {value!r}")
|
|
277
|
+
if value < 1:
|
|
278
|
+
raise ValueError(f"{command} needs '{name}' of at least 1, got {value}")
|
|
279
|
+
return value
|
|
280
|
+
|
|
281
|
+
|
|
282
|
+
def _evenly_spaced_indices(total, count):
|
|
283
|
+
"""`count` frame indices spaced evenly across [0, total - 1], inclusive
|
|
284
|
+
of both ends. Rounding can coincide two spacings on one index in a short
|
|
285
|
+
clip; those collapse rather than repeating the same frame as a tile."""
|
|
286
|
+
if count == 1:
|
|
287
|
+
return [0]
|
|
288
|
+
raw = numpy.linspace(0, total - 1, num=count)
|
|
289
|
+
seen = []
|
|
290
|
+
for value in raw.round().astype(int).tolist():
|
|
291
|
+
if not seen or seen[-1] != value:
|
|
292
|
+
seen.append(value)
|
|
293
|
+
return seen
|
|
294
|
+
|
|
295
|
+
|
|
296
|
+
def _default_columns(count):
|
|
297
|
+
"""A grid biased wide: rows no more than columns, columns >= sqrt(count)."""
|
|
298
|
+
rows = math.isqrt(count) or 1
|
|
299
|
+
return math.ceil(count / rows)
|
|
300
|
+
|
|
301
|
+
|
|
302
|
+
def _grid_tile(frame, index, fps, tile_width, label):
|
|
303
|
+
tile_height = max(1, round(frame.height * tile_width / frame.width))
|
|
304
|
+
tile = frame.resize((tile_width, tile_height), Image.LANCZOS).convert("RGB")
|
|
305
|
+
if not label:
|
|
306
|
+
return tile
|
|
307
|
+
|
|
308
|
+
from PIL import ImageDraw, ImageFont
|
|
309
|
+
|
|
310
|
+
text = _format_timestamp(index, fps) if fps else f"#{index}"
|
|
311
|
+
draw = ImageDraw.Draw(tile)
|
|
312
|
+
font_size = max(10, tile_width // 16)
|
|
313
|
+
try:
|
|
314
|
+
font = ImageFont.truetype("Arial", font_size)
|
|
315
|
+
except (IOError, OSError):
|
|
316
|
+
font = ImageFont.load_default(size=font_size)
|
|
317
|
+
draw.text(
|
|
318
|
+
(4, 4), text, font=font, fill="white", stroke_width=2, stroke_fill="black"
|
|
319
|
+
)
|
|
320
|
+
return tile
|
|
321
|
+
|
|
322
|
+
|
|
323
|
+
def _format_timestamp(index, fps):
|
|
324
|
+
seconds = index / fps
|
|
325
|
+
minutes, remainder = divmod(seconds, 60)
|
|
326
|
+
return f"{int(minutes):02d}:{remainder:04.1f}"
|
|
327
|
+
|
|
328
|
+
|
|
329
|
+
def _compose_grid(tiles, columns):
|
|
330
|
+
tile_width, tile_height = tiles[0].size
|
|
331
|
+
rows = math.ceil(len(tiles) / columns)
|
|
332
|
+
grid = Image.new("RGB", (columns * tile_width, rows * tile_height), (0, 0, 0))
|
|
333
|
+
for position, tile in enumerate(tiles):
|
|
334
|
+
row, col = divmod(position, columns)
|
|
335
|
+
grid.paste(tile, (col * tile_width, row * tile_height))
|
|
336
|
+
return grid
|
|
337
|
+
|
|
338
|
+
|
|
339
|
+
def is_video(value):
|
|
340
|
+
"""Whether a value is a run of frames rather than one image.
|
|
341
|
+
|
|
342
|
+
An AudioVideo, a 4-dim frame array or tensor, or a list of frames. A
|
|
343
|
+
single PIL image, a 3-dim array (one frame) and anything else is not.
|
|
344
|
+
"""
|
|
345
|
+
if isinstance(value, AudioVideo):
|
|
346
|
+
return True
|
|
347
|
+
if isinstance(value, list):
|
|
348
|
+
return len(value) > 0 and all(_is_frame(item) for item in value)
|
|
349
|
+
if isinstance(value, numpy.ndarray) or torch.is_tensor(value):
|
|
350
|
+
return value.ndim == 4 or (value.ndim == 5 and value.shape[0] == 1)
|
|
351
|
+
return False
|
|
352
|
+
|
|
353
|
+
|
|
354
|
+
def _frames_of(video):
|
|
355
|
+
"""Unwrap containers until an indexable run of frames remains."""
|
|
356
|
+
if isinstance(video, AudioVideo):
|
|
357
|
+
return _frames_of(video.frames)
|
|
358
|
+
|
|
359
|
+
# A bare still - e.g. a {"media_type": "image", ...} reference fetch_video
|
|
360
|
+
# now loads as a plain PIL image (#443) - is a one-frame video, the same
|
|
361
|
+
# accommodation loop_frames already made for itself with _is_frame
|
|
362
|
+
if isinstance(video, Image.Image):
|
|
363
|
+
return [video]
|
|
364
|
+
|
|
365
|
+
if isinstance(video, list):
|
|
366
|
+
# A one-video batch - [[frame, ...]] or [ndarray] - unwraps to the video;
|
|
367
|
+
# a single-frame video - [frame] - is already the frames
|
|
368
|
+
if len(video) == 1 and not _is_frame(video[0]):
|
|
369
|
+
return _frames_of(video[0])
|
|
370
|
+
return video
|
|
371
|
+
|
|
372
|
+
if isinstance(video, numpy.ndarray):
|
|
373
|
+
if video.ndim == 3: # a lone frame
|
|
374
|
+
return video[numpy.newaxis, ...]
|
|
375
|
+
if video.ndim == 5 and video.shape[0] == 1: # a one-video batch
|
|
376
|
+
return video[0]
|
|
377
|
+
return video
|
|
378
|
+
|
|
379
|
+
if torch.is_tensor(video):
|
|
380
|
+
tensor = video.detach().cpu()
|
|
381
|
+
if tensor.ndim == 5 and tensor.shape[0] == 1: # a one-video batch
|
|
382
|
+
tensor = tensor[0]
|
|
383
|
+
if tensor.ndim == 3: # a lone frame
|
|
384
|
+
tensor = tensor.unsqueeze(0)
|
|
385
|
+
return tensor
|
|
386
|
+
|
|
387
|
+
raise TypeError(f"Cannot extract frames from a {type(video).__name__}")
|
|
388
|
+
|
|
389
|
+
|
|
390
|
+
def _is_frame(item):
|
|
391
|
+
"""A single image: PIL, or a 3-dim array/tensor (height, width, channels)."""
|
|
392
|
+
if isinstance(item, Image.Image):
|
|
393
|
+
return True
|
|
394
|
+
if isinstance(item, numpy.ndarray) or torch.is_tensor(item):
|
|
395
|
+
return item.ndim == 3
|
|
396
|
+
return False
|
|
397
|
+
|
|
398
|
+
|
|
399
|
+
def _to_pil(frame):
|
|
400
|
+
"""Convert one frame to a PIL image; PIL frames pass through untouched."""
|
|
401
|
+
if isinstance(frame, Image.Image):
|
|
402
|
+
return frame
|
|
403
|
+
|
|
404
|
+
if torch.is_tensor(frame):
|
|
405
|
+
frame = frame.detach().cpu().float().numpy()
|
|
406
|
+
|
|
407
|
+
if isinstance(frame, numpy.ndarray):
|
|
408
|
+
if frame.ndim != 3:
|
|
409
|
+
raise ValueError(f"A frame must have 3 dimensions, got {frame.ndim}")
|
|
410
|
+
|
|
411
|
+
# Channels-first (C, H, W) -> channels-last, the layout PIL expects
|
|
412
|
+
if frame.shape[0] in (1, 3, 4) and frame.shape[-1] not in (1, 3, 4):
|
|
413
|
+
frame = numpy.moveaxis(frame, 0, -1)
|
|
414
|
+
|
|
415
|
+
if frame.dtype != numpy.uint8:
|
|
416
|
+
# Float frames are [0, 1] - diffusers' np output convention
|
|
417
|
+
frame = (numpy.clip(frame, 0.0, 1.0) * 255).round().astype(numpy.uint8)
|
|
418
|
+
|
|
419
|
+
if frame.shape[-1] == 1: # grayscale
|
|
420
|
+
frame = frame[..., 0]
|
|
421
|
+
|
|
422
|
+
return Image.fromarray(frame)
|
|
423
|
+
|
|
424
|
+
raise TypeError(f"Cannot convert a {type(frame).__name__} to an image")
|
|
425
|
+
|
|
426
|
+
|
|
427
|
+
class FrameList(list):
|
|
428
|
+
"""The frames of a video file, carrying the rate the file plays at and
|
|
429
|
+
the shot boundaries its run recorded, if any.
|
|
430
|
+
|
|
431
|
+
`load_video` answers a plain list of images, which is what every
|
|
432
|
+
pipeline argument and every task wants - and which says nothing about
|
|
433
|
+
how fast those frames are meant to run. A step handed a 24 fps file
|
|
434
|
+
then wrote it back at `result.fps`'s default of 8, three times long,
|
|
435
|
+
with its soundtrack finishing a third of the way in and nothing said
|
|
436
|
+
about it (#104, the file-loading half of #84). A list subclass keeps
|
|
437
|
+
every consumer working unchanged while `getattr(video, "fps", None)` -
|
|
438
|
+
the question AudioVideo, concat_videos and interpolate_frames already
|
|
439
|
+
ask - gets a real answer.
|
|
440
|
+
|
|
441
|
+
`shots` is the same idea for the boundaries `dw.runs.shots_beside`
|
|
442
|
+
finds beside the file: a video loaded from an `asset:`/`output:` path
|
|
443
|
+
carried no way to answer `getattr(video, "shots", None)`, so
|
|
444
|
+
`pair_audio` had nothing to remeasure even though the file's own
|
|
445
|
+
manifest (or its kept-asset sidecar) already held them (#398).
|
|
446
|
+
"""
|
|
447
|
+
|
|
448
|
+
def __init__(self, frames, fps=None, shots=None):
|
|
449
|
+
super().__init__(frames)
|
|
450
|
+
self.fps = fps
|
|
451
|
+
self.shots = shots
|
|
452
|
+
|
|
453
|
+
|
|
454
|
+
def file_fps(path):
|
|
455
|
+
"""The rate a video file declares, or None - a container that will not
|
|
456
|
+
open, carries no video stream or states no rate is a rate we do not
|
|
457
|
+
know, never an error: the caller is loading frames it has already read.
|
|
458
|
+
"""
|
|
459
|
+
try:
|
|
460
|
+
import av
|
|
461
|
+
|
|
462
|
+
with av.open(path) as container:
|
|
463
|
+
stream = container.streams.video[0] if container.streams.video else None
|
|
464
|
+
return (
|
|
465
|
+
float(stream.average_rate) if stream and stream.average_rate else None
|
|
466
|
+
)
|
|
467
|
+
except Exception as e:
|
|
468
|
+
logger.debug(f"No frame rate for {path}: {e}")
|
|
469
|
+
return None
|
|
470
|
+
|
|
471
|
+
|
|
472
|
+
def load_audio_video(location, base_dir=None):
|
|
473
|
+
"""Load a video file - frames and the audio muxed with them - as an AudioVideo.
|
|
474
|
+
|
|
475
|
+
`load_video` reads frames only, so a file written by an earlier run comes
|
|
476
|
+
back silent. Reading both streams here is what lets a step join videos that
|
|
477
|
+
are already on disk - the shots of an earlier run picked back up by name -
|
|
478
|
+
without dropping the audio those runs generated alongside them.
|
|
479
|
+
|
|
480
|
+
Args:
|
|
481
|
+
location: Local path, or an http(s) URL, of a video file
|
|
482
|
+
base_dir: Directory a relative path is resolved against
|
|
483
|
+
|
|
484
|
+
Returns:
|
|
485
|
+
An AudioVideo holding the frames as PIL images and, when the file
|
|
486
|
+
carries an audio stream, its waveform as a (channels, samples) float32
|
|
487
|
+
array with the stream's sample rate. A local file also carries the
|
|
488
|
+
shots its own run manifest recorded for it (`shots_beside`), so a
|
|
489
|
+
join of a file that is itself an earlier join's output can see the
|
|
490
|
+
seams inside it (#399); a URL carries none.
|
|
491
|
+
"""
|
|
492
|
+
from ..security import ALLOWED_VIDEO_EXTENSIONS, validate_file_extension
|
|
493
|
+
from ..locations import safe_get, validate_media_path
|
|
494
|
+
|
|
495
|
+
if _URL_SCHEME.match(location):
|
|
496
|
+
import io
|
|
497
|
+
|
|
498
|
+
# Any other scheme - ftp:, file:, data: - is refused here rather than
|
|
499
|
+
# falling through to be read as a relative path that happens to
|
|
500
|
+
# contain a colon. An http(s) one still has to name a host outside
|
|
501
|
+
# this deployment (dw/locations.py), and so does every redirect
|
|
502
|
+
logger.debug(f"Downloading video from {location}")
|
|
503
|
+
response = safe_get(location, "a video argument", timeout=300)
|
|
504
|
+
handle = io.BytesIO(response.content)
|
|
505
|
+
return _decode_audio_video(handle)
|
|
506
|
+
|
|
507
|
+
validated_path = validate_media_path(location, base_dir, "a video argument")
|
|
508
|
+
validate_file_extension(validated_path, ALLOWED_VIDEO_EXTENSIONS)
|
|
509
|
+
logger.debug(f"Reading video from {validated_path}")
|
|
510
|
+
video = _decode_audio_video(validated_path)
|
|
511
|
+
from ..runs import shots_beside
|
|
512
|
+
|
|
513
|
+
video.shots = shots_beside(validated_path)
|
|
514
|
+
return video
|
|
515
|
+
|
|
516
|
+
|
|
517
|
+
def _decode_audio_video(handle):
|
|
518
|
+
"""Decode a path or file object's video and audio streams in one pass."""
|
|
519
|
+
import av
|
|
520
|
+
from av.audio.resampler import AudioResampler
|
|
521
|
+
|
|
522
|
+
frames = []
|
|
523
|
+
chunks = []
|
|
524
|
+
sample_rate = None
|
|
525
|
+
|
|
526
|
+
with av.open(handle) as container:
|
|
527
|
+
video_stream = container.streams.video[0]
|
|
528
|
+
frame_rate = (
|
|
529
|
+
float(video_stream.average_rate) if video_stream.average_rate else None
|
|
530
|
+
)
|
|
531
|
+
streams = [video_stream]
|
|
532
|
+
if container.streams.audio:
|
|
533
|
+
audio_stream = container.streams.audio[0]
|
|
534
|
+
streams.append(audio_stream)
|
|
535
|
+
sample_rate = audio_stream.rate
|
|
536
|
+
# Planar float is the layout AudioVideo carries: (channels, samples)
|
|
537
|
+
resampler = AudioResampler(format="fltp")
|
|
538
|
+
|
|
539
|
+
for frame in container.decode(*streams):
|
|
540
|
+
if isinstance(frame, av.VideoFrame):
|
|
541
|
+
frames.append(Image.fromarray(frame.to_ndarray(format="rgb24")))
|
|
542
|
+
else:
|
|
543
|
+
chunks.extend(f.to_ndarray() for f in resampler.resample(frame))
|
|
544
|
+
|
|
545
|
+
if sample_rate is not None:
|
|
546
|
+
chunks.extend(f.to_ndarray() for f in resampler.resample(None))
|
|
547
|
+
|
|
548
|
+
audio = numpy.concatenate(chunks, axis=1).astype(numpy.float32) if chunks else None
|
|
549
|
+
if audio is not None and frame_rate:
|
|
550
|
+
audio = _fit_audio_to_frames(audio, len(frames), frame_rate, sample_rate)
|
|
551
|
+
logger.debug(
|
|
552
|
+
f"Decoded {len(frames)} frames and "
|
|
553
|
+
f"{audio.shape[1] if audio is not None else 0} audio samples"
|
|
554
|
+
)
|
|
555
|
+
# The file's own rate travels with it: a step that joins videos read
|
|
556
|
+
# from disk knows what to write them back at without being told (#84).
|
|
557
|
+
# A file carries no shots - the manifest that recorded them is the run's,
|
|
558
|
+
# not the file's
|
|
559
|
+
return AudioVideo(
|
|
560
|
+
frames, audio, sample_rate if audio is not None else None, fps=frame_rate
|
|
561
|
+
)
|
|
562
|
+
|
|
563
|
+
|
|
564
|
+
# How far a decoded track may be off the frames' own duration and still be
|
|
565
|
+
# treated as codec padding rather than a track of its own length. AAC codes
|
|
566
|
+
# 1024 samples at a time, so a file's audio runs up to one such block long -
|
|
567
|
+
# a hundredth of a second, which accumulates into visible lip-sync drift once
|
|
568
|
+
# a dozen shots are joined end to end
|
|
569
|
+
AUDIO_FIT_TOLERANCE_SECONDS = 0.25
|
|
570
|
+
|
|
571
|
+
|
|
572
|
+
def _fit_audio_to_frames(audio, frame_count, frame_rate, sample_rate):
|
|
573
|
+
"""Trim or pad a decoded track to exactly the frames' own duration.
|
|
574
|
+
|
|
575
|
+
Only when the difference is codec padding. A track that genuinely runs to
|
|
576
|
+
a different length than the picture - a song laid over a short clip - is
|
|
577
|
+
left alone.
|
|
578
|
+
|
|
579
|
+
`audio` may be a numpy array (the decode path) or a torch tensor still on
|
|
580
|
+
its generating device (an in-memory pipeline output, #197) - the pad and
|
|
581
|
+
trim below keep whichever type and device it arrived with rather than
|
|
582
|
+
forcing a host round trip the caller may not want yet.
|
|
583
|
+
"""
|
|
584
|
+
axis = _sample_axis(audio)
|
|
585
|
+
if axis is None:
|
|
586
|
+
return audio
|
|
587
|
+
|
|
588
|
+
expected = round(frame_count / frame_rate * sample_rate)
|
|
589
|
+
difference = audio.shape[axis] - expected
|
|
590
|
+
if difference == 0 or abs(difference) > AUDIO_FIT_TOLERANCE_SECONDS * sample_rate:
|
|
591
|
+
return audio
|
|
592
|
+
|
|
593
|
+
logger.debug(
|
|
594
|
+
f"Fitting decoded audio to {frame_count} frames ({difference:+} samples)"
|
|
595
|
+
)
|
|
596
|
+
if difference > 0:
|
|
597
|
+
trim = [slice(None)] * audio.ndim
|
|
598
|
+
trim[axis] = slice(None, expected)
|
|
599
|
+
return audio[tuple(trim)]
|
|
600
|
+
if isinstance(audio, torch.Tensor):
|
|
601
|
+
# torch.nn.functional.pad takes its pairs from the last axis backwards
|
|
602
|
+
padding = [0, 0] * audio.ndim
|
|
603
|
+
padding[2 * (audio.ndim - 1 - axis) + 1] = -difference
|
|
604
|
+
return torch.nn.functional.pad(audio, padding)
|
|
605
|
+
widths = [(0, 0)] * audio.ndim
|
|
606
|
+
widths[axis] = (0, -difference)
|
|
607
|
+
return numpy.pad(audio, widths)
|
|
608
|
+
|
|
609
|
+
|
|
610
|
+
def _sample_axis(audio):
|
|
611
|
+
"""The axis a waveform's samples run along, or None if it has no such axis.
|
|
612
|
+
|
|
613
|
+
Not a fixed index: a generated track arrives in any of the layouts
|
|
614
|
+
_as_stereo reads - (channels, samples), (samples, channels), or a bare
|
|
615
|
+
(samples,) - and a mono one written (samples,) or (samples, 1) used to
|
|
616
|
+
reach shape[1] here and either raise IndexError or fit the wrong axis
|
|
617
|
+
into a silent no-op. Channels are few and samples are many, so the
|
|
618
|
+
longer axis is the sample axis.
|
|
619
|
+
"""
|
|
620
|
+
if audio.ndim == 1:
|
|
621
|
+
return 0
|
|
622
|
+
if audio.ndim != 2:
|
|
623
|
+
return None
|
|
624
|
+
return 0 if audio.shape[0] > audio.shape[1] else 1
|
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
|